mirror of
https://github.com/strands-agents/harness-sdk.git
synced 2026-10-02 02:44:48 +08:00
Skip location for non bedrock model providers (#1602)
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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,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.
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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(
|
||||
{
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
]
|
||||
|
||||
@@ -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"])
|
||||
|
||||
@@ -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)
|
||||
@@ -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
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user