Refactor file parsing and loading in DirectoryComponent

This commit is contained in:
Gabriel Luiz Freitas Almeida 2024-03-08 17:06:41 -03:00
commit 865bc02c94
2 changed files with 74 additions and 15 deletions

View file

@ -1,9 +1,28 @@
import json
import xml.etree.ElementTree as ET
from concurrent import futures from concurrent import futures
from pathlib import Path from pathlib import Path
from typing import List, Optional, Text from typing import Callable, List, Optional, Text
import yaml
from langflow.schema.schema import Record from langflow.schema.schema import Record
# Types of files that can be read simply by file.read()
# and have 100% to be completely readable
TEXT_FILE_TYPES = [
"txt",
"md",
"mdx",
"csv",
"json",
"yaml",
"yml",
"xml",
"html",
"htm",
]
def is_hidden(path: Path) -> bool: def is_hidden(path: Path) -> bool:
return path.name.startswith(".") return path.name.startswith(".")
@ -11,10 +30,10 @@ def is_hidden(path: Path) -> bool:
def retrieve_file_paths( def retrieve_file_paths(
path: str, path: str,
types: List[str],
load_hidden: bool, load_hidden: bool,
recursive: bool, recursive: bool,
depth: int, depth: int,
types: List[str] = TEXT_FILE_TYPES,
) -> List[str]: ) -> List[str]:
path_obj = Path(path) path_obj = Path(path)
if not path_obj.exists() or not path_obj.is_dir(): if not path_obj.exists() or not path_obj.is_dir():
@ -35,12 +54,14 @@ def retrieve_file_paths(
glob = "**/*" if recursive else "*" glob = "**/*" if recursive else "*"
paths = walk_level(path_obj, depth) if depth else path_obj.glob(glob) paths = walk_level(path_obj, depth) if depth else path_obj.glob(glob)
file_paths = [Text(p) for p in paths if p.is_file() and match_types(p) and is_not_hidden(p)] file_paths = [
Text(p) for p in paths if p.is_file() and match_types(p) and is_not_hidden(p)
]
return file_paths return file_paths
def parse_file_to_record(file_path: str, silent_errors: bool) -> Optional[Record]: def partition_file_to_record(file_path: str, silent_errors: bool) -> Optional[Record]:
# Use the partition function to load the file # Use the partition function to load the file
from unstructured.partition.auto import partition # type: ignore from unstructured.partition.auto import partition # type: ignore
@ -59,6 +80,33 @@ def parse_file_to_record(file_path: str, silent_errors: bool) -> Optional[Record
return record return record
def read_text_file(file_path: str) -> str:
with open(file_path, "r") as f:
return f.read()
def parse_text_file_to_record(file_path: str, silent_errors: bool) -> Optional[Record]:
try:
text = read_text_file(file_path)
# if file is json, yaml, or xml, we can parse it
if file_path.endswith(".json"):
text = json.loads(text)
elif file_path.endswith(".yaml") or file_path.endswith(".yml"):
text = yaml.safe_load(text)
elif file_path.endswith(".xml"):
text = ET.fromstring(text)
except Exception as e:
if not silent_errors:
raise ValueError(f"Error loading file {file_path}: {e}") from e
return None
record = Record(data={"file_path": file_path, "text": text})
return record
def get_elements( def get_elements(
file_paths: List[str], file_paths: List[str],
silent_errors: bool, silent_errors: bool,
@ -68,15 +116,23 @@ def get_elements(
if use_multithreading: if use_multithreading:
records = parallel_load_records(file_paths, silent_errors, max_concurrency) records = parallel_load_records(file_paths, silent_errors, max_concurrency)
else: else:
records = [parse_file_to_record(file_path, silent_errors) for file_path in file_paths] records = [
partition_file_to_record(file_path, silent_errors)
for file_path in file_paths
]
records = list(filter(None, records)) records = list(filter(None, records))
return records return records
def parallel_load_records(file_paths: List[str], silent_errors: bool, max_concurrency: int) -> List[Optional[Record]]: def parallel_load_records(
file_paths: List[str],
silent_errors: bool,
max_concurrency: int,
load_function: Callable = parse_text_file_to_record,
) -> List[Optional[Record]]:
with futures.ThreadPoolExecutor(max_workers=max_concurrency) as executor: with futures.ThreadPoolExecutor(max_workers=max_concurrency) as executor:
loaded_files = executor.map( loaded_files = executor.map(
lambda file_path: parse_file_to_record(file_path, silent_errors), lambda file_path: load_function(file_path, silent_errors),
file_paths, file_paths,
) )
# loaded_files is an iterator, so we need to convert it to a list # loaded_files is an iterator, so we need to convert it to a list

View file

@ -3,7 +3,7 @@ from typing import Any, Dict, List, Optional
from langflow import CustomComponent from langflow import CustomComponent
from langflow.base.data.utils import ( from langflow.base.data.utils import (
parallel_load_records, parallel_load_records,
parse_file_to_record, parse_text_file_to_record,
retrieve_file_paths, retrieve_file_paths,
) )
from langflow.schema import Record from langflow.schema import Record
@ -11,7 +11,7 @@ from langflow.schema import Record
class DirectoryComponent(CustomComponent): class DirectoryComponent(CustomComponent):
display_name = "Directory" display_name = "Directory"
description = "Load files from a directory." description = "Load Text Files from a Directory and Convert Them to Records."
def build_config(self) -> Dict[str, Any]: def build_config(self) -> Dict[str, Any]:
return { return {
@ -46,7 +46,6 @@ class DirectoryComponent(CustomComponent):
def build( def build(
self, self,
path: str, path: str,
types: Optional[List[str]] = None,
depth: int = 0, depth: int = 0,
max_concurrency: int = 2, max_concurrency: int = 2,
load_hidden: bool = False, load_hidden: bool = False,
@ -54,16 +53,20 @@ class DirectoryComponent(CustomComponent):
silent_errors: bool = False, silent_errors: bool = False,
use_multithreading: bool = True, use_multithreading: bool = True,
) -> List[Optional[Record]]: ) -> List[Optional[Record]]:
if types is None:
types = []
resolved_path = self.resolve_path(path) resolved_path = self.resolve_path(path)
file_paths = retrieve_file_paths(resolved_path, types, load_hidden, recursive, depth) file_paths = retrieve_file_paths(resolved_path, load_hidden, recursive, depth)
loaded_records = [] loaded_records = []
if use_multithreading: if use_multithreading:
loaded_records = parallel_load_records(file_paths, silent_errors, max_concurrency) loaded_records = parallel_load_records(
file_paths, silent_errors, max_concurrency
)
else: else:
loaded_records = [parse_file_to_record(file_path, silent_errors) for file_path in file_paths] loaded_records = [
parse_text_file_to_record(file_path, silent_errors)
for file_path in file_paths
]
loaded_records = list(filter(None, loaded_records)) loaded_records = list(filter(None, loaded_records))
self.status = loaded_records self.status = loaded_records
return loaded_records return loaded_records