✨ (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:
parent
fcdef335ca
commit
318e028c58
3 changed files with 50 additions and 5 deletions
|
|
@ -738,7 +738,9 @@ class Graph:
|
||||||
# Check the cache for the vertex
|
# Check the cache for the vertex
|
||||||
cached_result = await chat_service.get_cache(key=vertex.id)
|
cached_result = await chat_service.get_cache(key=vertex.id)
|
||||||
if isinstance(cached_result, CacheMiss):
|
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)
|
await chat_service.set_cache(key=vertex.id, data=vertex)
|
||||||
else:
|
else:
|
||||||
cached_vertex = cached_result["result"]
|
cached_vertex = cached_result["result"]
|
||||||
|
|
@ -752,7 +754,9 @@ class Graph:
|
||||||
vertex.result.used_frozen_result = True
|
vertex.result.used_frozen_result = True
|
||||||
|
|
||||||
else:
|
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:
|
if vertex.result is not None:
|
||||||
params = f"{vertex._built_object_repr()}{params}"
|
params = f"{vertex._built_object_repr()}{params}"
|
||||||
|
|
|
||||||
|
|
@ -4,7 +4,8 @@ from typing import Literal, Optional, cast
|
||||||
|
|
||||||
from langchain_core.documents import Document
|
from langchain_core.documents import Document
|
||||||
from langchain_core.messages import AIMessage, BaseMessage, HumanMessage
|
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):
|
class Record(BaseModel):
|
||||||
|
|
@ -29,6 +30,11 @@ class Record(BaseModel):
|
||||||
values["data"][key] = values[key]
|
values["data"][key] = values[key]
|
||||||
return values
|
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):
|
def get_text(self):
|
||||||
"""
|
"""
|
||||||
Retrieves the text value from the data dictionary.
|
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)
|
text = self.data.pop(self.text_key, self.default_value)
|
||||||
return Document(page_content=text, metadata=self.data)
|
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.
|
Converts the Record to a BaseMessage.
|
||||||
|
|
||||||
|
|
@ -122,6 +130,30 @@ class Record(BaseModel):
|
||||||
return HumanMessage(content=text)
|
return HumanMessage(content=text)
|
||||||
return AIMessage(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):
|
def __getattr__(self, key):
|
||||||
"""
|
"""
|
||||||
Allows attribute-like access to the data dictionary.
|
Allows attribute-like access to the data dictionary.
|
||||||
|
|
@ -169,8 +201,14 @@ class Record(BaseModel):
|
||||||
|
|
||||||
def __str__(self) -> str:
|
def __str__(self) -> str:
|
||||||
# return a JSON string representation of the Record atributes
|
# 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"
|
INPUT_FIELD_NAME = "input_value"
|
||||||
|
|
|
||||||
|
|
@ -23,6 +23,9 @@ class CacheMiss:
|
||||||
def __repr__(self):
|
def __repr__(self):
|
||||||
return "<CACHE_MISS>"
|
return "<CACHE_MISS>"
|
||||||
|
|
||||||
|
def __bool__(self):
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
def create_cache_folder(func):
|
def create_cache_folder(func):
|
||||||
def wrapper(*args, **kwargs):
|
def wrapper(*args, **kwargs):
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue