fix: remove use of ImagePromptTemplate in image handling and adds image utils (#4467)

* remove poetry.lock

* fix: update langchain-core dependency version to 0.3.15

* feat: add functions to convert images to base64 and create data URLs

* refactor: simplify image URL handling by replacing ImagePromptTemplate with create_data_url function

* Fix image URL structure in data schema to use nested dictionary format

* Add unit tests for Data schema message conversion with text and images

* test: add unit tests for image utility functions to validate base64 conversion and data URL creation

* Refactor image URL generation to use `create_data_url` utility function instead of `ImagePromptTemplate`

* Add unit tests for message handling and image processing in schema module

- Introduce fixtures for temporary cache directory and sample image creation.
- Add tests for message creation from human and AI text.
- Implement tests for messages with single and multiple images.
- Include tests for invalid image paths and messages without sender.
- Add message serialization and conversion tests.
- Ensure cleanup of cache directory after tests.

* Use platformdirs to determine cache directory paths in test_schema_message.py
This commit is contained in:
Gabriel Luiz Freitas Almeida 2024-11-08 13:24:19 -03:00 • committed by GitHub
commit 98274633db
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
9 changed files with 889 additions and 12297 deletions

View file

@ -2,19 +2,16 @@ import copy
import json
from datetime import datetime
from decimal import Decimal
from typing import TYPE_CHECKING, cast
from typing import cast
from uuid import UUID
from langchain_core.documents import Document
from langchain_core.messages import AIMessage, BaseMessage, HumanMessage
from langchain_core.prompts.image import ImagePromptTemplate
from loguru import logger
from pydantic import BaseModel, model_serializer, model_validator
from langflow.utils.constants import MESSAGE_SENDER_AI, MESSAGE_SENDER_USER
if TYPE_CHECKING:
from langchain_core.prompt_values import ImagePromptValue
from langflow.utils.image import create_data_url
class Data(BaseModel):
@ -140,11 +137,8 @@ class Data(BaseModel):
if files:
contents = [{"type": "text", "text": text}]
for file_path in files:
image_template = ImagePromptTemplate()
image_prompt_value: ImagePromptValue = image_template.invoke(
input={"path": file_path}, config={"callbacks": self.get_langchain_callbacks()}
)
contents.append({"type": "image_url", "image_url": image_prompt_value.image_url})
image_url = create_data_url(file_path)
contents.append({"type": "image_url", "image_url": {"url": image_url}})
human_message = HumanMessage(content=contents)
else:
human_message = HumanMessage(

View file

@ -5,14 +5,13 @@ import re
import traceback
from collections.abc import AsyncIterator, Iterator
from datetime import datetime, timezone
from typing import TYPE_CHECKING, Annotated, Any, Literal
from typing import Annotated, Any, Literal
from uuid import UUID
from fastapi.encoders import jsonable_encoder
from langchain_core.load import load
from langchain_core.messages import AIMessage, BaseMessage, HumanMessage, SystemMessage
from langchain_core.prompts import BaseChatPromptTemplate, ChatPromptTemplate, PromptTemplate
from langchain_core.prompts.image import ImagePromptTemplate
from loguru import logger
from pydantic import BaseModel, ConfigDict, Field, ValidationError, field_serializer, field_validator
@ -29,9 +28,7 @@ from langflow.utils.constants import (
MESSAGE_SENDER_NAME_USER,
MESSAGE_SENDER_USER,
)
if TYPE_CHECKING:
from langchain_core.prompt_values import ImagePromptValue
from langflow.utils.image import create_data_url
class Message(Data):
@ -203,9 +200,8 @@ class Message(Data):
if isinstance(file, Image):
content_dicts.append(file.to_content_dict())
else:
image_template = ImagePromptTemplate()
image_prompt_value: ImagePromptValue = image_template.invoke(input={"path": file})
content_dicts.append({"type": "image_url", "image_url": image_prompt_value.image_url})
image_url = create_data_url(file)
content_dicts.append({"type": "image_url", "image_url": {"url": image_url}})
return content_dicts
def load_lc_prompt(self):

View file

@ -0,0 +1,69 @@
import base64
import mimetypes
from pathlib import Path
def convert_image_to_base64(image_path: str | Path) -> str:
"""Convert an image file to a base64 encoded string.
Args:
image_path (str | Path): Path to the image file.
Returns:
str: Base64 encoded string representation of the image.
Raises:
FileNotFoundError: If the image file does not exist.
IOError: If there's an error reading the image file.
ValueError: If the image path is empty or invalid.
"""
if not image_path:
msg = "Image path cannot be empty"
raise ValueError(msg)
image_path = Path(image_path)
if not image_path.exists():
msg = f"Image file not found: {image_path}"
raise FileNotFoundError(msg)
if not image_path.is_file():
msg = f"Path is not a file: {image_path}"
raise ValueError(msg)
try:
with image_path.open("rb") as image_file:
return base64.b64encode(image_file.read()).decode("utf-8")
except OSError as e:
msg = f"Error reading image file: {e}"
raise OSError(msg) from e
def create_data_url(image_path: str | Path, mime_type: str | None = None) -> str:
"""Create a data URL from an image file.
Args:
image_path (str | Path): Path to the image file.
mime_type (Optional[str], optional): MIME type of the image.
If None, it will be guessed from the file extension.
Returns:
str: Data URL containing the base64 encoded image.
Raises:
FileNotFoundError: If the image file does not exist.
IOError: If there's an error reading the image file.
ValueError: If the image path is empty or invalid.
"""
if not mime_type:
mime_type = mimetypes.guess_type(str(image_path))[0]
if not mime_type:
msg = f"Could not determine MIME type for: {image_path}"
raise ValueError(msg)
try:
base64_data = convert_image_to_base64(image_path)
except (OSError, FileNotFoundError, ValueError) as e:
msg = f"Failed to create data URL: {e}"
raise type(e)(msg) from e
return f"data:{mime_type};base64,{base64_data}"

View file

@ -111,7 +111,7 @@ dependencies = [
"uvicorn>=0.30.0",
"gunicorn>=22.0.0",
"langchain~=0.3.3",
"langchain-core~=0.3.10",
"langchain-core~=0.3.15",
"langchainhub~=0.1.15",
"sqlmodel==0.0.18",
"loguru>=0.7.1",