Skip location for non bedrock model providers (#1602)

This commit is contained in:
Nick Clegg
2026-01-30 11:40:42 -05:00
committed by GitHub
parent f2c35a593d
commit ab51706c2a
20 changed files with 708 additions and 41 deletions
+21
View File
@@ -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
+6 -1
View File
@@ -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:
+1 -1
View File
@@ -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:
+18 -9
View File
@@ -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.
+12 -6
View File
@@ -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
+12 -6
View File
@@ -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(
{
+6 -1
View File
@@ -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):
+11 -7
View File
@@ -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
+12 -6
View File
@@ -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
]
+13 -3
View File
@@ -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"])
+67
View File
@@ -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)
+67
View File
@@ -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
+1 -1
View File
@@ -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):
+64
View File
@@ -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
+67
View File
@@ -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
+69
View File
@@ -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
+63
View File
@@ -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
+66
View File
@@ -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
+65
View File
@@ -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
+67
View File
@@ -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