Refactor RecursiveCharacterTextSplitterComponent to use build_loader_repr_from_records

This commit is contained in:
Gabriel Luiz Freitas Almeida 2024-03-06 15:32:50 -03:00
commit f735e50fd2
2 changed files with 59 additions and 28 deletions

View file

@ -5,13 +5,15 @@ from langchain_core.documents import Document
from langflow import CustomComponent from langflow import CustomComponent
from langflow.schema import Record from langflow.schema import Record
from langflow.utils.util import build_loader_repr_from_documents from langflow.utils.util import build_loader_repr_from_records
class RecursiveCharacterTextSplitterComponent(CustomComponent): class RecursiveCharacterTextSplitterComponent(CustomComponent):
display_name: str = "Recursive Character Text Splitter" display_name: str = "Recursive Character Text Splitter"
description: str = "Split text into chunks of a specified length." description: str = "Split text into chunks of a specified length."
documentation: str = "https://docs.langflow.org/components/text-splitters#recursivecharactertextsplitter" documentation: str = (
"https://docs.langflow.org/components/text-splitters#recursivecharactertextsplitter"
)
def build_config(self): def build_config(self):
return { return {
@ -84,5 +86,6 @@ class RecursiveCharacterTextSplitterComponent(CustomComponent):
else: else:
documents.append(_input) documents.append(_input)
docs = splitter.split_documents(documents) docs = splitter.split_documents(documents)
self.repr_value = build_loader_repr_from_documents(docs) records = self.to_records(docs)
return self.to_records(docs) self.repr_value = build_loader_repr_from_records(records)
return records

View file

@ -5,8 +5,8 @@ from functools import wraps
from typing import Any, Dict, List, Optional, Union from typing import Any, Dict, List, Optional, Union
from docstring_parser import parse from docstring_parser import parse
from langchain_core.documents import Document
from langflow.schema.schema import Record
from langflow.template.frontend_node.constants import FORCE_SHOW_FIELDS from langflow.template.frontend_node.constants import FORCE_SHOW_FIELDS
from langflow.utils import constants from langflow.utils import constants
@ -15,8 +15,12 @@ def remove_ansi_escape_codes(text):
return re.sub(r"\x1b\[[0-9;]*[a-zA-Z]", "", text) return re.sub(r"\x1b\[[0-9;]*[a-zA-Z]", "", text)
def build_template_from_function(name: str, type_to_loader_dict: Dict, add_function: bool = False): def build_template_from_function(
classes = [item.__annotations__["return"].__name__ for item in type_to_loader_dict.values()] name: str, type_to_loader_dict: Dict, add_function: bool = False
):
classes = [
item.__annotations__["return"].__name__ for item in type_to_loader_dict.values()
]
# Raise error if name is not in chains # Raise error if name is not in chains
if name not in classes: if name not in classes:
@ -37,8 +41,10 @@ def build_template_from_function(name: str, type_to_loader_dict: Dict, add_funct
for name_, value_ in value.__repr_args__(): for name_, value_ in value.__repr_args__():
if name_ == "default_factory": if name_ == "default_factory":
try: try:
variables[class_field_items]["default"] = get_default_factory( variables[class_field_items]["default"] = (
module=_class.__base__.__module__, function=value_ get_default_factory(
module=_class.__base__.__module__, function=value_
)
) )
except Exception: except Exception:
variables[class_field_items]["default"] = None variables[class_field_items]["default"] = None
@ -46,7 +52,9 @@ def build_template_from_function(name: str, type_to_loader_dict: Dict, add_funct
variables[class_field_items][name_] = value_ variables[class_field_items][name_] = value_
variables[class_field_items]["placeholder"] = ( variables[class_field_items]["placeholder"] = (
docs.params[class_field_items] if class_field_items in docs.params else "" docs.params[class_field_items]
if class_field_items in docs.params
else ""
) )
# Adding function to base classes to allow # Adding function to base classes to allow
# the output to be a function # the output to be a function
@ -61,7 +69,9 @@ def build_template_from_function(name: str, type_to_loader_dict: Dict, add_funct
} }
def build_template_from_class(name: str, type_to_cls_dict: Dict, add_function: bool = False): def build_template_from_class(
name: str, type_to_cls_dict: Dict, add_function: bool = False
):
classes = [item.__name__ for item in type_to_cls_dict.values()] classes = [item.__name__ for item in type_to_cls_dict.values()]
# Raise error if name is not in chains # Raise error if name is not in chains
@ -85,9 +95,11 @@ def build_template_from_class(name: str, type_to_cls_dict: Dict, add_function: b
for name_, value_ in value.__repr_args__(): for name_, value_ in value.__repr_args__():
if name_ == "default_factory": if name_ == "default_factory":
try: try:
variables[class_field_items]["default"] = get_default_factory( variables[class_field_items]["default"] = (
module=_class.__base__.__module__, get_default_factory(
function=value_, module=_class.__base__.__module__,
function=value_,
)
) )
except Exception: except Exception:
variables[class_field_items]["default"] = None variables[class_field_items]["default"] = None
@ -95,7 +107,9 @@ def build_template_from_class(name: str, type_to_cls_dict: Dict, add_function: b
variables[class_field_items][name_] = value_ variables[class_field_items][name_] = value_
variables[class_field_items]["placeholder"] = ( variables[class_field_items]["placeholder"] = (
docs.params[class_field_items] if class_field_items in docs.params else "" docs.params[class_field_items]
if class_field_items in docs.params
else ""
) )
base_classes = get_base_classes(_class) base_classes = get_base_classes(_class)
# Adding function to base classes to allow # Adding function to base classes to allow
@ -127,7 +141,9 @@ def build_template_from_method(
# Check if the method exists in this class # Check if the method exists in this class
if not hasattr(_class, method_name): if not hasattr(_class, method_name):
raise ValueError(f"Method {method_name} not found in class {class_name}") raise ValueError(
f"Method {method_name} not found in class {class_name}"
)
# Get the method # Get the method
method = getattr(_class, method_name) method = getattr(_class, method_name)
@ -146,8 +162,14 @@ def build_template_from_method(
"_type": _type, "_type": _type,
**{ **{
name: { name: {
"default": (param.default if param.default != param.empty else None), "default": (
"type": (param.annotation if param.annotation != param.empty else None), param.default if param.default != param.empty else None
),
"type": (
param.annotation
if param.annotation != param.empty
else None
),
"required": param.default == param.empty, "required": param.default == param.empty,
} }
for name, param in params.items() for name, param in params.items()
@ -234,7 +256,9 @@ def sync_to_async(func):
return async_wrapper return async_wrapper
def format_dict(dictionary: Dict[str, Any], class_name: Optional[str] = None) -> Dict[str, Any]: def format_dict(
dictionary: Dict[str, Any], class_name: Optional[str] = None
) -> Dict[str, Any]:
""" """
Formats a dictionary by removing certain keys and modifying the Formats a dictionary by removing certain keys and modifying the
values of other keys. values of other keys.
@ -320,7 +344,9 @@ def check_list_type(_type: str, value: Dict[str, Any]) -> str:
The modified type string. The modified type string.
""" """
if any(list_type in _type for list_type in ["List", "Sequence", "Set"]): if any(list_type in _type for list_type in ["List", "Sequence", "Set"]):
_type = _type.replace("List[", "").replace("Sequence[", "").replace("Set[", "")[:-1] _type = (
_type.replace("List[", "").replace("Sequence[", "").replace("Set[", "")[:-1]
)
value["list"] = True value["list"] = True
else: else:
value["list"] = False value["list"] = False
@ -423,7 +449,9 @@ def set_headers_value(value: Dict[str, Any]) -> None:
value["value"] = """{"Authorization": "Bearer <token>"}""" value["value"] = """{"Authorization": "Bearer <token>"}"""
def add_options_to_field(value: Dict[str, Any], class_name: Optional[str], key: str) -> None: def add_options_to_field(
value: Dict[str, Any], class_name: Optional[str], key: str
) -> None:
""" """
Adds options to the field based on the class name and key. Adds options to the field based on the class name and key.
""" """
@ -440,10 +468,10 @@ def add_options_to_field(value: Dict[str, Any], class_name: Optional[str], key:
value["value"] = options_map[class_name][0] value["value"] = options_map[class_name][0]
def build_loader_repr_from_documents(documents: List[Document]) -> str: def build_loader_repr_from_records(records: List[Record]) -> str:
if documents: if records:
avg_length = sum(len(doc.page_content) for doc in documents) / len(documents) avg_length = sum(len(doc.page_content) for doc in records) / len(records)
return f"""{len(documents)} documents return f"""{len(records)} records
\nAvg. Document Length (characters): {int(avg_length)} \nAvg. Record Length (characters): {int(avg_length)}
Documents: {documents[:3]}...""" Records: {records[:3]}..."""
return "0 documents" return "0 records"