diff --git a/src/strands/models/_validation.py b/src/strands/models/_validation.py index 1e82bca73..9d4d8b178 100644 --- a/src/strands/models/_validation.py +++ b/src/strands/models/_validation.py @@ -6,6 +6,7 @@ from typing import Any from typing_extensions import get_type_hints +from ..types.content import ContentBlock from ..types.tools import ToolChoice @@ -41,3 +42,23 @@ def warn_on_tool_choice_not_supported(tool_choice: ToolChoice | None) -> None: "A ToolChoice was provided to this provider but is not supported and will be ignored", stacklevel=4, ) + + +def _has_location_source(content: ContentBlock) -> bool: + """Check if a content block contains a location source. + + Providers need to explicitly define an implementation to support content locations. + + Args: + content: Content block to check. + + Returns: + True if the content block contains an location source, False otherwise. + """ + if "image" in content: + return "location" in content["image"].get("source", {}) + if "document" in content: + return "location" in content["document"].get("source", {}) + if "video" in content: + return "location" in content["video"].get("source", {}) + return False diff --git a/src/strands/models/anthropic.py b/src/strands/models/anthropic.py index 535c820ee..b5f6fcf91 100644 --- a/src/strands/models/anthropic.py +++ b/src/strands/models/anthropic.py @@ -20,7 +20,7 @@ from ..types.content import ContentBlock, Messages from ..types.exceptions import ContextWindowOverflowException, ModelThrottledException from ..types.streaming import StreamEvent from ..types.tools import ToolChoice, ToolChoiceToolDict, ToolSpec -from ._validation import validate_config_keys +from ._validation import _has_location_source, validate_config_keys from .model import Model logger = logging.getLogger(__name__) @@ -189,6 +189,11 @@ class AnthropicModel(Model): formatted_contents[-1]["cache_control"] = {"type": "ephemeral"} continue + # Check for location sources in image, document, or video content + if _has_location_source(content): + logger.warning("Location sources are not supported by Anthropic | skipping content block") + continue + formatted_contents.append(self._format_request_message_content(content)) if formatted_contents: diff --git a/src/strands/models/bedrock.py b/src/strands/models/bedrock.py index b053b70fb..596936e6f 100644 --- a/src/strands/models/bedrock.py +++ b/src/strands/models/bedrock.py @@ -472,7 +472,7 @@ class BedrockModel(Model): formatted_document_s3["bucketOwner"] = s3_location["bucketOwner"] return {"s3Location": formatted_document_s3} else: - logger.warning("Non s3 location sources are not supported by Bedrock, skipping content block") + logger.warning("Non s3 location sources are not supported by Bedrock | skipping content block") return None def _format_request_message_content(self, content: ContentBlock) -> dict[str, Any] | None: diff --git a/src/strands/models/gemini.py b/src/strands/models/gemini.py index 192a363d3..6a6535999 100644 --- a/src/strands/models/gemini.py +++ b/src/strands/models/gemini.py @@ -18,7 +18,7 @@ from ..types.content import ContentBlock, Messages from ..types.exceptions import ContextWindowOverflowException, ModelThrottledException from ..types.streaming import StreamEvent from ..types.tools import ToolChoice, ToolSpec -from ._validation import validate_config_keys +from ._validation import _has_location_source, validate_config_keys from .model import Model logger = logging.getLogger(__name__) @@ -229,15 +229,24 @@ class GeminiModel(Model): # available in tool result blocks, hence the mapping. tool_use_id_to_name: dict[str, str] = {} - return [ - genai.types.Content( - parts=[ - self._format_request_content_part(content, tool_use_id_to_name) for content in message["content"] - ], - role="user" if message["role"] == "user" else "model", + contents = [] + for message in messages: + parts = [] + for content in message["content"]: + # Check for location sources and skip with warning + if _has_location_source(content): + logger.warning("Location sources are not supported by Gemini | skipping content block") + continue + parts.append(self._format_request_content_part(content, tool_use_id_to_name)) + + contents.append( + genai.types.Content( + parts=parts, + role="user" if message["role"] == "user" else "model", + ) ) - for message in messages - ] + + return contents def _format_request_tools(self, tool_specs: list[ToolSpec] | None) -> list[genai.types.Tool | Any]: """Format tool specs into Gemini tools. diff --git a/src/strands/models/llamaapi.py b/src/strands/models/llamaapi.py index ce0367bf5..b1ed4563a 100644 --- a/src/strands/models/llamaapi.py +++ b/src/strands/models/llamaapi.py @@ -20,7 +20,7 @@ from ..types.content import ContentBlock, Messages from ..types.exceptions import ModelThrottledException from ..types.streaming import StreamEvent, Usage from ..types.tools import ToolChoice, ToolResult, ToolSpec, ToolUse -from ._validation import validate_config_keys, warn_on_tool_choice_not_supported +from ._validation import _has_location_source, validate_config_keys, warn_on_tool_choice_not_supported from .model import Model logger = logging.getLogger(__name__) @@ -176,12 +176,18 @@ class LlamaAPIModel(Model): for message in messages: contents = message["content"] + # Filter out location sources and unsupported block types + filtered_contents = [] + for content in contents: + if any(block_type in content for block_type in ["toolResult", "toolUse"]): + continue + if _has_location_source(content): + logger.warning("Location sources are not supported by LlamaAPI | skipping content block") + continue + filtered_contents.append(content) + formatted_contents: list[dict[str, Any]] | dict[str, Any] | str = "" - formatted_contents = [ - self._format_request_message_content(content) - for content in contents - if not any(block_type in content for block_type in ["toolResult", "toolUse"]) - ] + formatted_contents = [self._format_request_message_content(content) for content in filtered_contents] formatted_tool_calls = [ self._format_request_message_tool_call(content["toolUse"]) for content in contents diff --git a/src/strands/models/llamacpp.py b/src/strands/models/llamacpp.py index ca838f3d7..c52509816 100644 --- a/src/strands/models/llamacpp.py +++ b/src/strands/models/llamacpp.py @@ -30,7 +30,7 @@ from ..types.content import ContentBlock, Messages from ..types.exceptions import ContextWindowOverflowException, ModelThrottledException from ..types.streaming import StreamEvent from ..types.tools import ToolChoice, ToolSpec -from ._validation import validate_config_keys, warn_on_tool_choice_not_supported +from ._validation import _has_location_source, validate_config_keys, warn_on_tool_choice_not_supported from .model import Model logger = logging.getLogger(__name__) @@ -299,11 +299,17 @@ class LlamaCppModel(Model): for message in messages: contents = message["content"] - formatted_contents = [ - self._format_message_content(content) - for content in contents - if not any(block_type in content for block_type in ["toolResult", "toolUse"]) - ] + # Filter out location sources and unsupported block types + filtered_contents = [] + for content in contents: + if any(block_type in content for block_type in ["toolResult", "toolUse"]): + continue + if _has_location_source(content): + logger.warning("Location sources are not supported by llama.cpp | skipping content block") + continue + filtered_contents.append(content) + + formatted_contents = [self._format_message_content(content) for content in filtered_contents] formatted_tool_calls = [ self._format_tool_call( { diff --git a/src/strands/models/mistral.py b/src/strands/models/mistral.py index 4ec77ccfe..504e81c92 100644 --- a/src/strands/models/mistral.py +++ b/src/strands/models/mistral.py @@ -17,7 +17,7 @@ from ..types.content import ContentBlock, Messages from ..types.exceptions import ModelThrottledException from ..types.streaming import StopReason, StreamEvent from ..types.tools import ToolChoice, ToolResult, ToolSpec, ToolUse -from ._validation import validate_config_keys, warn_on_tool_choice_not_supported +from ._validation import _has_location_source, validate_config_keys, warn_on_tool_choice_not_supported from .model import Model logger = logging.getLogger(__name__) @@ -212,6 +212,11 @@ class MistralModel(Model): tool_messages: list[dict[str, Any]] = [] for content in contents: + # Check for location sources and skip with warning + if _has_location_source(content): + logger.warning("Location sources are not supported by Mistral | skipping content block") + continue + if "text" in content: formatted_content = self._format_request_message_content(content) if isinstance(formatted_content, str): diff --git a/src/strands/models/ollama.py b/src/strands/models/ollama.py index 8d72aa534..68aba59d4 100644 --- a/src/strands/models/ollama.py +++ b/src/strands/models/ollama.py @@ -15,7 +15,7 @@ from typing_extensions import TypedDict, Unpack, override from ..types.content import ContentBlock, Messages from ..types.streaming import StopReason, StreamEvent from ..types.tools import ToolChoice, ToolSpec -from ._validation import validate_config_keys, warn_on_tool_choice_not_supported +from ._validation import _has_location_source, validate_config_keys, warn_on_tool_choice_not_supported from .model import Model logger = logging.getLogger(__name__) @@ -160,12 +160,16 @@ class OllamaModel(Model): """ system_message = [{"role": "system", "content": system_prompt}] if system_prompt else [] - return system_message + [ - formatted_message - for message in messages - for content in message["content"] - for formatted_message in self._format_request_message_contents(message["role"], content) - ] + formatted_messages = [] + for message in messages: + for content in message["content"]: + # Check for location sources and skip with warning + if _has_location_source(content): + logger.warning("Location sources are not supported by Ollama | skipping content block") + continue + formatted_messages.extend(self._format_request_message_contents(message["role"], content)) + + return system_message + formatted_messages def format_request( self, messages: Messages, tool_specs: list[ToolSpec] | None = None, system_prompt: str | None = None diff --git a/src/strands/models/openai.py b/src/strands/models/openai.py index d9266212b..51e98c8c2 100644 --- a/src/strands/models/openai.py +++ b/src/strands/models/openai.py @@ -20,7 +20,7 @@ from ..types.content import ContentBlock, Messages, SystemContentBlock from ..types.exceptions import ContextWindowOverflowException, ModelThrottledException from ..types.streaming import StreamEvent from ..types.tools import ToolChoice, ToolResult, ToolSpec, ToolUse -from ._validation import validate_config_keys +from ._validation import _has_location_source, validate_config_keys from .model import Model logger = logging.getLogger(__name__) @@ -338,11 +338,17 @@ class OpenAIModel(Model): "reasoningContent is not supported in multi-turn conversations with the Chat Completions API." ) - formatted_contents = [ - cls.format_request_message_content(content) - for content in contents - if not any(block_type in content for block_type in ["toolResult", "toolUse", "reasoningContent"]) - ] + # Filter out content blocks that shouldn't be formatted + filtered_contents = [] + for content in contents: + if any(block_type in content for block_type in ["toolResult", "toolUse", "reasoningContent"]): + continue + if _has_location_source(content): + logger.warning("Location sources are not supported by OpenAI | skipping content block") + continue + filtered_contents.append(content) + + formatted_contents = [cls.format_request_message_content(content) for content in filtered_contents] formatted_tool_calls = [ cls.format_request_message_tool_call(content["toolUse"]) for content in contents if "toolUse" in content ] diff --git a/src/strands/models/writer.py b/src/strands/models/writer.py index f306d649b..94774b363 100644 --- a/src/strands/models/writer.py +++ b/src/strands/models/writer.py @@ -18,7 +18,7 @@ from ..types.content import ContentBlock, Messages from ..types.exceptions import ModelThrottledException from ..types.streaming import StreamEvent from ..types.tools import ToolChoice, ToolResult, ToolSpec, ToolUse -from ._validation import validate_config_keys, warn_on_tool_choice_not_supported +from ._validation import _has_location_source, validate_config_keys, warn_on_tool_choice_not_supported from .model import Model logger = logging.getLogger(__name__) @@ -218,11 +218,21 @@ class WriterModel(Model): for message in messages: contents = message["content"] + # Filter out location sources + filtered_contents = [] + for content in contents: + if _has_location_source(content): + logger.warning("Location sources are not supported by Writer | skipping content block") + continue + filtered_contents.append(content) + # Only palmyra V5 support multiple content. Other models support only '{"content": "text_content"}' if self.get_config().get("model_id", "") == "palmyra-x5": - formatted_contents: str | list[dict[str, Any]] = self._format_request_message_contents_vision(contents) + formatted_contents: str | list[dict[str, Any]] = self._format_request_message_contents_vision( + filtered_contents + ) else: - formatted_contents = self._format_request_message_contents(contents) + formatted_contents = self._format_request_message_contents(filtered_contents) formatted_tool_calls = [ self._format_request_message_tool_call(content["toolUse"]) diff --git a/tests/strands/models/test__validation.py b/tests/strands/models/test__validation.py new file mode 100644 index 000000000..e8a451494 --- /dev/null +++ b/tests/strands/models/test__validation.py @@ -0,0 +1,67 @@ +"""Tests for model validation helper functions.""" + +from strands.models._validation import _has_location_source + + +class TestHasLocationSource: + """Tests for _has_location_source helper function.""" + + def test_image_with_location_source(self): + """Test detection of location source in image content.""" + content = {"image": {"source": {"location": {"type": "s3", "uri": "s3://bucket/key"}}}} + assert _has_location_source(content) + + def test_image_with_bytes_source(self): + """Test that bytes source is not detected as location.""" + content = {"image": {"source": {"bytes": b"data"}}} + assert not _has_location_source(content) + + def test_document_with_location_source(self): + """Test detection of location source in document content.""" + content = {"document": {"source": {"location": {"type": "s3", "uri": "s3://bucket/key"}}}} + assert _has_location_source(content) + + def test_document_with_bytes_source(self): + """Test that bytes source is not detected as location.""" + content = {"document": {"source": {"bytes": b"data"}}} + assert not _has_location_source(content) + + def test_video_with_location_source(self): + """Test detection of location source in video content.""" + content = {"video": {"source": {"location": {"type": "s3", "uri": "s3://bucket/key"}}}} + assert _has_location_source(content) + + def test_video_with_bytes_source(self): + """Test that bytes source is not detected as location.""" + content = {"video": {"source": {"bytes": b"data"}}} + assert not _has_location_source(content) + + def test_text_content(self): + """Test that text content is not detected as location source.""" + content = {"text": "hello"} + assert not _has_location_source(content) + + def test_tool_use_content(self): + """Test that toolUse content is not detected as location source.""" + content = {"toolUse": {"name": "test", "input": {}, "toolUseId": "123"}} + assert not _has_location_source(content) + + def test_tool_result_content(self): + """Test that toolResult content is not detected as location source.""" + content = {"toolResult": {"toolUseId": "123", "content": [{"text": "result"}]}} + assert not _has_location_source(content) + + def test_image_without_source(self): + """Test that image without source is not detected as location.""" + content = {"image": {"format": "png"}} + assert not _has_location_source(content) + + def test_document_without_source(self): + """Test that document without source is not detected as location.""" + content = {"document": {"format": "pdf", "name": "test.pdf"}} + assert not _has_location_source(content) + + def test_video_without_source(self): + """Test that video without source is not detected as location.""" + content = {"video": {"format": "mp4"}} + assert not _has_location_source(content) diff --git a/tests/strands/models/test_anthropic.py b/tests/strands/models/test_anthropic.py index 74bbb8d45..c5aff8062 100644 --- a/tests/strands/models/test_anthropic.py +++ b/tests/strands/models/test_anthropic.py @@ -1,3 +1,4 @@ +import logging import unittest.mock import anthropic @@ -866,3 +867,69 @@ def test_tool_choice_none_no_warning(model, messages, captured_warnings): model.format_request(messages, tool_choice=None) assert len(captured_warnings) == 0 + + +def test_format_request_filters_s3_source_image(model, model_id, max_tokens, caplog): + """Test that images with Location sources are filtered out with warning.""" + caplog.set_level(logging.WARNING, logger="strands.models.anthropic") + + messages = [ + { + "role": "user", + "content": [ + {"text": "look at this image"}, + { + "image": { + "format": "png", + "source": {"location": {"type": "s3", "uri": "s3://my-bucket/image.png"}}, + }, + }, + ], + }, + ] + + tru_request = model.format_request(messages) + + # Image with S3 source should be filtered, text should remain + exp_messages = [ + {"role": "user", "content": [{"type": "text", "text": "look at this image"}]}, + ] + assert tru_request["messages"] == exp_messages + assert "Location sources are not supported by Anthropic" in caplog.text + + +def test_format_request_filters_location_source_document(model, model_id, max_tokens, caplog): + """Test that documents with Location sources are filtered out with warning.""" + caplog.set_level(logging.WARNING, logger="strands.models.anthropic") + + messages = [ + { + "role": "user", + "content": [ + {"text": "analyze this document"}, + { + "document": { + "format": "pdf", + "name": "report.pdf", + "source": {"location": {"type": "s3", "uri": "s3://my-bucket/report.pdf"}}, + }, + }, + { + "document": { + "format": "pdf", + "name": "report.pdf", + "source": {"location": {"type": "s3", "uri": "s3://my-bucket/report.pdf"}}, + }, + }, + ], + }, + ] + + tru_request = model.format_request(messages) + + # Document with S3 source should be filtered, text should remain + exp_messages = [ + {"role": "user", "content": [{"type": "text", "text": "analyze this document"}]}, + ] + assert tru_request["messages"] == exp_messages + assert "Location sources are not supported by Anthropic" in caplog.text diff --git a/tests/strands/models/test_bedrock.py b/tests/strands/models/test_bedrock.py index 761434258..aac791214 100644 --- a/tests/strands/models/test_bedrock.py +++ b/tests/strands/models/test_bedrock.py @@ -1924,7 +1924,7 @@ def test_format_request_unsupported_location(model, caplog): formatted_request = model._format_request(messages) assert len(formatted_request["messages"][0]["content"]) == 1 - assert "Non s3 location sources are not supported by Bedrock, skipping content block" in caplog.text + assert "Non s3 location sources are not supported by Bedrock | skipping content block" in caplog.text def test_format_request_video_s3_location(model, model_id): diff --git a/tests/strands/models/test_gemini.py b/tests/strands/models/test_gemini.py index 86ab2fea5..d62c5a7c8 100644 --- a/tests/strands/models/test_gemini.py +++ b/tests/strands/models/test_gemini.py @@ -934,3 +934,67 @@ def test_init_with_both_client_and_client_args_raises_error(): with pytest.raises(ValueError, match="Only one of 'client' or 'client_args' should be provided"): GeminiModel(client=mock_client, client_args={"api_key": "test"}, model_id="test-model") + + +def test_format_request_filters_s3_source_image(model, caplog): + """Test that images with Location sources are filtered out with warning.""" + caplog.set_level(logging.WARNING, logger="strands.models.gemini") + + messages = [ + { + "role": "user", + "content": [ + {"text": "look at this image"}, + { + "image": { + "format": "png", + "source": {"location": {"type": "s3", "uri": "s3://my-bucket/image.png"}}, + }, + }, + ], + }, + ] + + request = model._format_request(messages, None, None, None) + + # Image with S3 source should be filtered, text should remain + formatted_content = request["contents"][0]["parts"] + assert len(formatted_content) == 1 + assert "text" in formatted_content[0] + assert "Location sources are not supported by Gemini" in caplog.text + + +def test_format_request_filters_location_source_document(model, caplog): + """Test that documents with Location sources are filtered out with warning.""" + caplog.set_level(logging.WARNING, logger="strands.models.gemini") + + messages = [ + { + "role": "user", + "content": [ + {"text": "analyze this document"}, + { + "document": { + "format": "pdf", + "name": "report.pdf", + "source": {"location": {"type": "s3", "uri": "s3://my-bucket/report.pdf"}}, + }, + }, + { + "document": { + "format": "pdf", + "name": "report.pdf", + "source": {"location": {"type": "s3", "uri": "s3://my-bucket/report.pdf"}}, + }, + }, + ], + }, + ] + + request = model._format_request(messages, None, None, None) + + # Document with S3 source should be filtered, text should remain + formatted_content = request["contents"][0]["parts"] + assert len(formatted_content) == 1 + assert "text" in formatted_content[0] + assert "Location sources are not supported by Gemini" in caplog.text diff --git a/tests/strands/models/test_llamaapi.py b/tests/strands/models/test_llamaapi.py index a6bbf5673..2bf12d055 100644 --- a/tests/strands/models/test_llamaapi.py +++ b/tests/strands/models/test_llamaapi.py @@ -1,4 +1,5 @@ # Copyright (c) Meta Platforms, Inc. and affiliates +import logging import unittest.mock import pytest @@ -414,3 +415,69 @@ async def test_tool_choice_none_no_warning(model, messages, captured_warnings, a await alist(response) assert len(captured_warnings) == 0 + + +def test_format_request_filters_s3_source_image(model, caplog): + """Test that images with Location sources are filtered out with warning.""" + caplog.set_level(logging.WARNING, logger="strands.models.llamaapi") + + messages = [ + { + "role": "user", + "content": [ + {"text": "look at this image"}, + { + "image": { + "format": "png", + "source": {"location": {"type": "s3", "uri": "s3://my-bucket/image.png"}}, + }, + }, + ], + }, + ] + + request = model.format_request(messages) + + # Image with S3 source should be filtered, text should remain + formatted_messages = request["messages"] + user_content = formatted_messages[0]["content"] + assert len(user_content) == 1 + assert user_content[0]["type"] == "text" + assert "Location sources are not supported by LlamaAPI" in caplog.text + + +def test_format_request_filters_location_source_document(model, caplog): + """Test that documents with Location sources are filtered out with warning.""" + caplog.set_level(logging.WARNING, logger="strands.models.llamaapi") + + messages = [ + { + "role": "user", + "content": [ + {"text": "analyze this document"}, + { + "document": { + "format": "pdf", + "name": "report.pdf", + "source": {"location": {"type": "s3", "uri": "s3://my-bucket/report.pdf"}}, + }, + }, + { + "document": { + "format": "pdf", + "name": "report.pdf", + "source": {"location": {"type": "s3", "uri": "s3://my-bucket/report.pdf"}}, + }, + }, + ], + }, + ] + + request = model.format_request(messages) + + # Document with S3 source should be filtered, text should remain + formatted_messages = request["messages"] + user_content = formatted_messages[0]["content"] + assert len(user_content) == 1 + assert user_content[0]["type"] == "text" + assert "Location sources are not supported by LlamaAPI" in caplog.text diff --git a/tests/strands/models/test_llamacpp.py b/tests/strands/models/test_llamacpp.py index e5b2614c0..fa784de5c 100644 --- a/tests/strands/models/test_llamacpp.py +++ b/tests/strands/models/test_llamacpp.py @@ -2,6 +2,7 @@ import base64 import json +import logging from unittest.mock import AsyncMock, patch import httpx @@ -637,3 +638,71 @@ def test_format_messages_with_mixed_content() -> None: assert result[0]["content"][2]["type"] == "image_url" assert "image_url" in result[0]["content"][2] assert result[0]["content"][2]["image_url"]["url"].startswith("data:image/jpeg;base64,") + + +def test_format_request_filters_s3_source_image(caplog) -> None: + """Test that images with Location sources are filtered out with warning.""" + model = LlamaCppModel() + caplog.set_level(logging.WARNING, logger="strands.models.llamacpp") + + messages = [ + { + "role": "user", + "content": [ + {"text": "look at this image"}, + { + "image": { + "format": "png", + "source": {"location": {"type": "s3", "uri": "s3://my-bucket/image.png"}}, + }, + }, + ], + }, + ] + + request = model._format_request(messages) + + # Image with S3 source should be filtered, text should remain + formatted_messages = request["messages"] + user_content = formatted_messages[0]["content"] + assert len(user_content) == 1 + assert user_content[0]["type"] == "text" + assert "Location sources are not supported by llama.cpp" in caplog.text + + +def test_format_request_filters_location_source_document(caplog) -> None: + """Test that documents with Location sources are filtered out with warning.""" + model = LlamaCppModel() + caplog.set_level(logging.WARNING, logger="strands.models.llamacpp") + + messages = [ + { + "role": "user", + "content": [ + {"text": "analyze this document"}, + { + "document": { + "format": "pdf", + "name": "report.pdf", + "source": {"location": {"type": "s3", "uri": "s3://my-bucket/report.pdf"}}, + }, + }, + { + "document": { + "format": "pdf", + "name": "report.pdf", + "source": {"location": {"type": "s3", "uri": "s3://my-bucket/report.pdf"}}, + }, + }, + ], + }, + ] + + request = model._format_request(messages) + + # Document with S3 source should be filtered, text should remain + formatted_messages = request["messages"] + user_content = formatted_messages[0]["content"] + assert len(user_content) == 1 + assert user_content[0]["type"] == "text" + assert "Location sources are not supported by llama.cpp" in caplog.text diff --git a/tests/strands/models/test_mistral.py b/tests/strands/models/test_mistral.py index 7808336f2..ad74bae89 100644 --- a/tests/strands/models/test_mistral.py +++ b/tests/strands/models/test_mistral.py @@ -1,3 +1,4 @@ +import logging import unittest.mock import pydantic @@ -592,3 +593,65 @@ def test_update_config_validation_warns_on_unknown_keys(model, captured_warnings assert len(captured_warnings) == 1 assert "Invalid configuration parameters" in str(captured_warnings[0].message) assert "wrong_param" in str(captured_warnings[0].message) + + +def test_format_request_filters_s3_source_image(model, caplog): + """Test that images with Location sources are filtered out with warning.""" + caplog.set_level(logging.WARNING, logger="strands.models.mistral") + + messages = [ + { + "role": "user", + "content": [ + {"text": "look at this image"}, + { + "image": { + "format": "png", + "source": {"location": {"type": "s3", "uri": "s3://my-bucket/image.png"}}, + }, + }, + ], + }, + ] + + formatted_messages = model._format_request_messages(messages) + + # Image with S3 source should be filtered, text should remain + user_content = formatted_messages[0]["content"] + assert user_content == "look at this image" + assert "Location sources are not supported by Mistral" in caplog.text + + +def test_format_request_filters_location_source_document(model, caplog): + """Test that documents with Location sources are filtered out with warning.""" + caplog.set_level(logging.WARNING, logger="strands.models.mistral") + + messages = [ + { + "role": "user", + "content": [ + {"text": "analyze this document"}, + { + "document": { + "format": "pdf", + "name": "report.pdf", + "source": {"location": {"type": "s3", "uri": "s3://my-bucket/report.pdf"}}, + }, + }, + { + "document": { + "format": "pdf", + "name": "report.pdf", + "source": {"location": {"type": "s3", "uri": "s3://my-bucket/report.pdf"}}, + }, + }, + ], + }, + ] + + formatted_messages = model._format_request_messages(messages) + + # Document with S3 source should be filtered, text should remain + user_content = formatted_messages[0]["content"] + assert user_content == "analyze this document" + assert "Location sources are not supported by Mistral" in caplog.text diff --git a/tests/strands/models/test_ollama.py b/tests/strands/models/test_ollama.py index 14db63a24..d17894028 100644 --- a/tests/strands/models/test_ollama.py +++ b/tests/strands/models/test_ollama.py @@ -1,4 +1,5 @@ import json +import logging import unittest.mock import pydantic @@ -559,3 +560,68 @@ def test_update_config_validation_warns_on_unknown_keys(model, captured_warnings assert len(captured_warnings) == 1 assert "Invalid configuration parameters" in str(captured_warnings[0].message) assert "wrong_param" in str(captured_warnings[0].message) + + +def test_format_request_filters_s3_source_image(model, caplog): + """Test that images with Location sources are filtered out with warning.""" + caplog.set_level(logging.WARNING, logger="strands.models.ollama") + + messages = [ + { + "role": "user", + "content": [ + {"text": "look at this image"}, + { + "image": { + "format": "png", + "source": {"location": {"type": "s3", "uri": "s3://my-bucket/image.png"}}, + }, + }, + ], + }, + ] + + request = model.format_request(messages) + + # Image with S3 source should be filtered, text should remain + formatted_messages = request["messages"] + user_message = formatted_messages[0] + assert user_message["content"] == "look at this image" + assert "images" not in user_message or user_message.get("images") == [] + assert "Location sources are not supported by Ollama" in caplog.text + + +def test_format_request_filters_location_source_document(model, caplog): + """Test that documents with Location sources are filtered out with warning.""" + caplog.set_level(logging.WARNING, logger="strands.models.ollama") + + messages = [ + { + "role": "user", + "content": [ + {"text": "analyze this document"}, + { + "document": { + "format": "pdf", + "name": "report.pdf", + "source": {"location": {"type": "s3", "uri": "s3://my-bucket/report.pdf"}}, + }, + }, + { + "document": { + "format": "pdf", + "name": "report.pdf", + "source": {"location": {"type": "s3", "uri": "s3://my-bucket/report.pdf"}}, + }, + }, + ], + }, + ] + + request = model.format_request(messages) + + # Document with S3 source should be filtered, text should remain + formatted_messages = request["messages"] + user_message = formatted_messages[0] + assert user_message["content"] == "analyze this document" + assert "Location sources are not supported by Ollama" in caplog.text diff --git a/tests/strands/models/test_openai.py b/tests/strands/models/test_openai.py index 7c1d18998..6eeb477d9 100644 --- a/tests/strands/models/test_openai.py +++ b/tests/strands/models/test_openai.py @@ -1,3 +1,4 @@ +import logging import unittest.mock import openai @@ -1246,3 +1247,67 @@ def test_init_with_both_client_and_client_args_raises_error(): with pytest.raises(ValueError, match="Only one of 'client' or 'client_args' should be provided"): OpenAIModel(client=mock_client, client_args={"api_key": "test"}, model_id="test-model") + + +def test_format_request_filters_s3_source_image(model, caplog): + """Test that images with Location sources are filtered out with warning.""" + caplog.set_level(logging.WARNING, logger="strands.models.openai") + + messages = [ + { + "role": "user", + "content": [ + {"text": "look at this image"}, + { + "image": { + "format": "png", + "source": {"location": {"type": "s3", "uri": "s3://my-bucket/image.png"}}, + }, + }, + ], + }, + ] + + request = model.format_request(messages) + + # Image with S3 source should be filtered, text should remain + formatted_content = request["messages"][0]["content"] + assert len(formatted_content) == 1 + assert formatted_content[0]["type"] == "text" + assert "Location sources are not supported by OpenAI" in caplog.text + + +def test_format_request_filters_location_source_document(model, caplog): + """Test that documents with Location sources are filtered out with warning.""" + caplog.set_level(logging.WARNING, logger="strands.models.openai") + + messages = [ + { + "role": "user", + "content": [ + {"text": "analyze this document"}, + { + "document": { + "format": "pdf", + "name": "report.pdf", + "source": {"location": {"type": "s3", "uri": "s3://my-bucket/report.pdf"}}, + }, + }, + { + "document": { + "format": "pdf", + "name": "report.pdf", + "source": {"location": {"type": "s3", "uri": "s3://my-bucket/report.pdf"}}, + }, + }, + ], + }, + ] + + request = model.format_request(messages) + + # Document with S3 source should be filtered, text should remain + formatted_content = request["messages"][0]["content"] + assert len(formatted_content) == 1 + assert formatted_content[0]["type"] == "text" + assert "Location sources are not supported by OpenAI" in caplog.text diff --git a/tests/strands/models/test_writer.py b/tests/strands/models/test_writer.py index 963904002..81745f412 100644 --- a/tests/strands/models/test_writer.py +++ b/tests/strands/models/test_writer.py @@ -1,3 +1,4 @@ +import logging import unittest.mock from typing import Any @@ -435,3 +436,69 @@ def test_update_config_validation_warns_on_unknown_keys(model, captured_warnings assert len(captured_warnings) == 1 assert "Invalid configuration parameters" in str(captured_warnings[0].message) assert "wrong_param" in str(captured_warnings[0].message) + + +def test_format_request_filters_s3_source_image(model, caplog): + """Test that images with Location sources are filtered out with warning.""" + caplog.set_level(logging.WARNING, logger="strands.models.writer") + + messages = [ + { + "role": "user", + "content": [ + {"text": "look at this image"}, + { + "image": { + "format": "png", + "source": {"location": {"type": "s3", "uri": "s3://my-bucket/image.png"}}, + }, + }, + ], + }, + ] + + request = model.format_request(messages) + + # Image with S3 source should be filtered, text should remain + formatted_messages = request["messages"] + user_content = formatted_messages[0]["content"] + assert len(user_content) == 1 + assert user_content[0]["type"] == "text" + assert "Location sources are not supported by Writer" in caplog.text + + +def test_format_request_filters_location_source_document(model, caplog): + """Test that documents with Location sources are filtered out with warning.""" + caplog.set_level(logging.WARNING, logger="strands.models.writer") + + messages = [ + { + "role": "user", + "content": [ + {"text": "analyze this document"}, + { + "document": { + "format": "pdf", + "name": "report.pdf", + "source": {"location": {"type": "s3", "uri": "s3://my-bucket/report.pdf"}}, + }, + }, + { + "document": { + "format": "pdf", + "name": "report.pdf", + "source": {"location": {"type": "s3", "uri": "s3://my-bucket/report.pdf"}}, + }, + }, + ], + }, + ] + + request = model.format_request(messages) + + # Document with S3 source should be filtered, text should remain + formatted_messages = request["messages"] + user_content = formatted_messages[0]["content"] + assert len(user_content) == 1 + assert user_content[0]["type"] == "text" + assert "Location sources are not supported by Writer" in caplog.text