From 318e028c58c9b17e31aa5d9396392664bd61856f Mon Sep 17 00:00:00 2001 From: ogabrielluiz Date: Thu, 6 Jun 2024 11:48:15 -0300 Subject: [PATCH] =?UTF-8?q?=E2=9C=A8=20(base.py):=20Add=20support=20for=20?= =?UTF-8?q?passing=20'files'=20parameter=20to=20vertex=20build=20method=20?= =?UTF-8?q?to=20handle=20file=20inputs?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ♻️ (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 --- src/backend/base/langflow/graph/graph/base.py | 8 +++- src/backend/base/langflow/schema/schema.py | 44 +++++++++++++++++-- .../base/langflow/services/cache/utils.py | 3 ++ 3 files changed, 50 insertions(+), 5 deletions(-) diff --git a/src/backend/base/langflow/graph/graph/base.py b/src/backend/base/langflow/graph/graph/base.py index b0e43557d..57ccc6a9d 100644 --- a/src/backend/base/langflow/graph/graph/base.py +++ b/src/backend/base/langflow/graph/graph/base.py @@ -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}" diff --git a/src/backend/base/langflow/schema/schema.py b/src/backend/base/langflow/schema/schema.py index 921bd65b2..624ba7dcb 100644 --- a/src/backend/base/langflow/schema/schema.py +++ b/src/backend/base/langflow/schema/schema.py @@ -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" diff --git a/src/backend/base/langflow/services/cache/utils.py b/src/backend/base/langflow/services/cache/utils.py index ff19836ef..a89963f56 100644 --- a/src/backend/base/langflow/services/cache/utils.py +++ b/src/backend/base/langflow/services/cache/utils.py @@ -23,6 +23,9 @@ class CacheMiss: def __repr__(self): return "" + def __bool__(self): + return False + def create_cache_folder(func): def wrapper(*args, **kwargs):