Files
AI/참고/instructor-main/tests/llm/test_anthropic/test_multimodal.py
2026-05-12 19:40:31 +09:00

242 lines
8.3 KiB
Python

import pytest
from instructor.processing.multimodal import Image, PDF, PDFWithCacheControl
import instructor
from pydantic import Field, BaseModel
from itertools import product
from .util import models, modes
import os
import base64
# Models that support PDF input (Claude 3.5+ only)
_PDF_CAPABLE = ["claude-3-5", "claude-3-7", "claude-haiku-4", "claude-sonnet-4"]
pdf_supported = any(cap in m for cap in _PDF_CAPABLE for m in models)
class ImageDescription(BaseModel):
objects: list[str] = Field(..., description="The objects in the image")
scene: str = Field(..., description="The scene of the image")
colors: list[str] = Field(..., description="The colors in the image")
image_url = "https://raw.githubusercontent.com/instructor-ai/instructor/main/tests/assets/image.jpg"
pdf_url = "https://raw.githubusercontent.com/instructor-ai/instructor/main/tests/assets/invoice.pdf"
curr_file = os.path.dirname(__file__)
pdf_path = os.path.join(curr_file, "../../assets/invoice.pdf")
pdf_base64 = base64.b64encode(open(pdf_path, "rb").read()).decode("utf-8")
pdf_base64_string = f"data:application/pdf;base64,{pdf_base64}"
@pytest.mark.parametrize("model, mode", product(models, modes))
def test_multimodal_image_description(model, mode):
client = instructor.from_provider(model, mode=mode)
response = client.chat.completions.create(
response_model=ImageDescription,
messages=[
{
"role": "system",
"content": "You are a helpful assistant that can describe images",
},
{
"role": "user",
"content": [
"What is this?",
Image.from_url(image_url),
],
},
],
temperature=1,
max_tokens=1000,
)
# Assertions to validate the response
assert isinstance(response, ImageDescription)
assert len(response.objects) > 0
assert response.scene != ""
assert len(response.colors) > 0
@pytest.mark.parametrize("model, mode", product(models, modes))
def test_multimodal_image_description_autodetect(model, mode):
client = instructor.from_provider(model, mode=mode)
response = client.chat.completions.create(
response_model=ImageDescription,
messages=[
{
"role": "system",
"content": "You are a helpful assistant that can describe images",
},
{
"role": "user",
"content": [
"What is this?",
image_url,
],
},
],
max_tokens=1000,
temperature=1,
autodetect_images=True,
)
# Assertions to validate the response
assert isinstance(response, ImageDescription)
assert len(response.objects) > 0
assert response.scene != ""
assert len(response.colors) > 0
# Additional assertions can be added based on expected content of the sample image
@pytest.mark.parametrize("model, mode", product(models, modes))
def test_multimodal_image_description_autodetect_image_params(model, mode):
client = instructor.from_provider(model, mode=mode)
response = client.chat.completions.create(
response_model=ImageDescription,
messages=[
{
"role": "system",
"content": "You are a helpful assistant that can describe images",
},
{
"role": "user",
"content": [
"What is this?",
{
"type": "image",
"source": image_url,
},
],
},
],
max_tokens=1000,
temperature=1,
autodetect_images=True,
)
# Assertions to validate the response
assert isinstance(response, ImageDescription)
assert len(response.objects) > 0
assert response.scene != ""
assert len(response.colors) > 0
# Additional assertions can be added based on expected content of the sample image
@pytest.mark.parametrize("model, mode", product(models, modes))
def test_multimodal_image_description_autodetect_image_params_cache(model, mode):
client = instructor.from_provider(model, mode=mode)
messages = client.chat.completions.create(
response_model=None,
messages=[
{
"role": "system",
"content": "You are a helpful assistant that can describe images and stuff",
},
{
"role": "user",
"content": [
"Describe these images",
# Large images to activate caching
{
"type": "image",
"source": "https://assets.entrepreneur.com/content/3x2/2000/20200429211042-GettyImages-1164615296.jpeg",
"cache_control": {"type": "ephemeral"},
},
{
"type": "image",
"source": "https://www.bigbear.com/imager/s3_us-west-1_amazonaws_com/big-bear/images/Scenic-Snow/89xVzXp1_00588cdef1e3d54756582b576359604b.jpeg",
"cache_control": {"type": "ephemeral"},
},
],
},
],
max_tokens=1000,
temperature=1,
autodetect_images=True,
)
# Cache tokens are non-deterministic (Anthropic may not always activate cache
# on first call or for small payloads). Just verify the fields are present.
assert hasattr(messages.usage, "cache_creation_input_tokens")
assert hasattr(messages.usage, "cache_read_input_tokens")
class LineItem(BaseModel):
name: str
price: int
quantity: int
class Receipt(BaseModel):
total: int
items: list[str]
@pytest.mark.skipif(not pdf_supported, reason="PDF input requires Claude 3.5+ models")
@pytest.mark.parametrize("pdf_source", [pdf_path, pdf_url, pdf_base64_string])
@pytest.mark.parametrize("model, mode", product(models, modes))
def test_multimodal_pdf_file(model, mode, pdf_source):
client = instructor.from_provider(model, mode=mode)
# Retry logic for flaky LLM responses
max_retries = 3
for attempt in range(max_retries):
response = client.chat.completions.create(
messages=[
{
"role": "system",
"content": "Extract the total and items from the invoice. Be precise and only extract the final total amount and list of item names. The total should be exactly 220.",
},
{
"role": "user",
"content": PDF.autodetect(pdf_source),
},
],
max_tokens=1000,
temperature=0, # Keep at 0 for consistent responses
autodetect_images=False,
response_model=Receipt,
)
if response.total == 220 and len(response.items) == 2:
break
elif attempt == max_retries - 1:
pytest.fail(
f"After {max_retries} attempts, got total={response.total}, items={response.items}, expected total=220, items=2"
)
assert response.total == 220
assert len(response.items) == 2
@pytest.mark.skipif(not pdf_supported, reason="PDF input requires Claude 3.5+ models")
@pytest.mark.parametrize("pdf_source", [pdf_path, pdf_url, pdf_base64_string])
@pytest.mark.parametrize("model, mode", product(models, modes))
def test_multimodal_pdf_file_with_cache_control(model, mode, pdf_source):
client = instructor.from_provider(model, mode=mode)
response, completion = client.chat.completions.create_with_completion(
messages=[
{
"role": "system",
"content": "Extract the total and items from the invoice",
},
{
"role": "user",
"content": PDFWithCacheControl.autodetect(pdf_source),
},
],
max_tokens=1000,
autodetect_images=False,
response_model=Receipt,
)
assert response.total == 220
# Cache tokens are non-deterministic. Just verify the fields exist.
assert hasattr(completion.usage, "cache_creation_input_tokens")
assert hasattr(completion.usage, "cache_read_input_tokens")
assert len(response.items) == 2