chore: Refactor loading.py to improve component instantiation and build process

This commit is contained in:
Gabriel Luiz Freitas Almeida 2024-06-19 18:54:19 -03:00
commit 9e4a414682

View file

@ -2,7 +2,7 @@ import inspect
import json import json
import os import os
import warnings import warnings
from typing import TYPE_CHECKING, Any from typing import TYPE_CHECKING, Any, Type
import orjson import orjson
from loguru import logger from loguru import logger
@ -12,9 +12,11 @@ from langflow.custom import Component, CustomComponent
from langflow.custom.eval import eval_custom_component_code from langflow.custom.eval import eval_custom_component_code
from langflow.schema import Data from langflow.schema import Data
from langflow.schema.artifact import get_artifact_type, post_process_raw from langflow.schema.artifact import get_artifact_type, post_process_raw
from langflow.services.deps import get_tracing_service
if TYPE_CHECKING: if TYPE_CHECKING:
from langflow.graph.vertex.base import Vertex from langflow.graph.vertex.base import Vertex
from langflow.services.tracing.service import TracingService
async def instantiate_class( async def instantiate_class(
@ -34,23 +36,40 @@ async def instantiate_class(
if not base_type: if not base_type:
raise ValueError("No base type provided for vertex") raise ValueError("No base type provided for vertex")
custom_component, build_results, artifacts = await build_component_and_get_results(
params=params,
vertex=vertex,
user_id=user_id,
tracing_service=get_tracing_service(),
fallback_to_env_vars=fallback_to_env_vars,
base_type=base_type,
)
return custom_component, build_results, artifacts
async def build_component_and_get_results(
params: dict,
vertex: "Vertex",
user_id: str,
tracing_service: "TracingService",
fallback_to_env_vars: bool = False,
base_type: str = "component",
):
params_copy = params.copy() params_copy = params.copy()
# Remove code from params # Remove code from params
class_object = eval_custom_component_code(params_copy.pop("code")) class_object: Type["CustomComponent" | "Component"] = eval_custom_component_code(params_copy.pop("code"))
custom_component = class_object( custom_component: "CustomComponent" | "Component" = class_object(
user_id=user_id, user_id=user_id, parameters=params_copy, vertex=vertex, tracing_service=tracing_service
parameters=params_copy,
vertex=vertex,
) )
params_copy = update_params_with_load_from_db_fields( params_copy = update_params_with_load_from_db_fields(
custom_component, params_copy, vertex.load_from_db_fields, fallback_to_env_vars custom_component, params_copy, vertex.load_from_db_fields, fallback_to_env_vars
) )
with warnings.catch_warnings(): with warnings.catch_warnings():
warnings.filterwarnings("ignore", category=PydanticDeprecatedSince20) warnings.filterwarnings("ignore", category=PydanticDeprecatedSince20)
if isinstance(custom_component, Component): if base_type == "custom_components":
return await build_component(params=params_copy, custom_component=custom_component, vertex=vertex)
elif isinstance(custom_component, CustomComponent):
return await build_custom_component(params=params_copy, custom_component=custom_component) return await build_custom_component(params=params_copy, custom_component=custom_component)
elif base_type == "component":
return await build_component(params=params_copy, custom_component=custom_component)
else: else:
raise ValueError(f"Base type {base_type} not found.") raise ValueError(f"Base type {base_type} not found.")
@ -121,11 +140,10 @@ def update_params_with_load_from_db_fields(
async def build_component( async def build_component(
params: dict, params: dict,
custom_component: "Component", custom_component: "Component",
vertex: "Vertex",
): ):
# Now set the params as attributes of the custom_component # Now set the params as attributes of the custom_component
custom_component.set_attributes(params) custom_component.set_attributes(params)
build_results, artifacts = await custom_component.build_results(vertex) build_results, artifacts = await custom_component.build_results()
return custom_component, build_results, artifacts return custom_component, build_results, artifacts