✨ (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
|
||||
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}"
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue