✨ (base.py): Add support for passing 'files' parameter to vertex build method to handle file inputs

♻️ (schema.py): Refactor code to include new method 'to_lc_messages' for converting Record to a list of BaseMessage
📝 (utils.py): Add __bool__ method to CacheMiss class to improve boolean evaluation
This commit is contained in:
ogabrielluiz 2024-06-06 11:48:15 -03:00
commit 318e028c58
3 changed files with 50 additions and 5 deletions

View file

@ -738,7 +738,9 @@ class Graph:
# Check the cache for the vertex
cached_result = await chat_service.get_cache(key=vertex.id)
if isinstance(cached_result, CacheMiss):
await vertex.build(user_id=user_id, inputs=inputs_dict, fallback_to_env_vars=fallback_to_env_vars)
await vertex.build(
user_id=user_id, inputs=inputs_dict, fallback_to_env_vars=fallback_to_env_vars, files=files
)
await chat_service.set_cache(key=vertex.id, data=vertex)
else:
cached_vertex = cached_result["result"]
@ -752,7 +754,9 @@ class Graph:
vertex.result.used_frozen_result = True
else:
await vertex.build(user_id=user_id, inputs=inputs_dict, fallback_to_env_vars=fallback_to_env_vars)
await vertex.build(
user_id=user_id, inputs=inputs_dict, fallback_to_env_vars=fallback_to_env_vars, files=files
)
if vertex.result is not None:
params = f"{vertex._built_object_repr()}{params}"

View file

@ -4,7 +4,8 @@ from typing import Literal, Optional, cast
from langchain_core.documents import Document
from langchain_core.messages import AIMessage, BaseMessage, HumanMessage
from pydantic import BaseModel, model_validator
from langchain_core.prompts.image import ImagePromptTemplate
from pydantic import BaseModel, model_serializer, model_validator
class Record(BaseModel):
@ -29,6 +30,11 @@ class Record(BaseModel):
values["data"][key] = values[key]
return values
@model_serializer(mode="json")
def serialize_model(cls, obj):
data = {k: v.to_json() if hasattr(v, "to_json") else v for k, v in obj.data.items()}
return data
def get_text(self):
"""
Retrieves the text value from the data dictionary.
@ -102,7 +108,9 @@ class Record(BaseModel):
text = self.data.pop(self.text_key, self.default_value)
return Document(page_content=text, metadata=self.data)
def to_lc_message(self) -> BaseMessage:
def to_lc_message(
self,
) -> BaseMessage:
"""
Converts the Record to a BaseMessage.
@ -122,6 +130,30 @@ class Record(BaseModel):
return HumanMessage(content=text)
return AIMessage(content=text)
def to_lc_messages(self):
"""
Converts the Record to a list of BaseMessage.
Returns:
list[BaseMessage]: The converted list of BaseMessage.
"""
if not all(key in self.data for key in ["text", "sender"]):
raise ValueError(f"Missing required keys ('text', 'sender') in Record: {self.data}")
sender = self.data.get("sender", "Machine")
text = self.data.get("text", "")
files = self.data.get("files", [])
if sender == "User":
if files:
human_messages = [HumanMessage(content=text)]
for base64_image in files:
image_template = ImagePromptTemplate()
human_message = image_template.invoke(url=f"data:image/png;base64,{base64_image}")
human_messages.append(human_message)
return human_messages
else:
return [HumanMessage(content=text)]
return [AIMessage(content=text)]
def __getattr__(self, key):
"""
Allows attribute-like access to the data dictionary.
@ -169,8 +201,14 @@ class Record(BaseModel):
def __str__(self) -> str:
# return a JSON string representation of the Record atributes
try:
data = {k: v.to_json() if hasattr(v, "to_json") else v for k, v in self.data.items()}
return json.dumps(data, indent=4)
except Exception:
return str(self.data)
return json.dumps(self.data)
def __contains__(self, key):
return key in self.data
INPUT_FIELD_NAME = "input_value"

View file

@ -23,6 +23,9 @@ class CacheMiss:
def __repr__(self):
return "<CACHE_MISS>"
def __bool__(self):
return False
def create_cache_folder(func):
def wrapper(*args, **kwargs):