Fix issues with user authentication and custom

component field config
This commit is contained in:
Gabriel Luiz Freitas Almeida 2023-11-03 13:57:14 -03:00
commit 364d36d194
6 changed files with 20 additions and 16 deletions

View file

@ -127,6 +127,7 @@ async def stream_build(
flow_id: str, flow_id: str,
chat_service: "ChatService" = Depends(get_chat_service), chat_service: "ChatService" = Depends(get_chat_service),
cache_service: "BaseCacheService" = Depends(get_cache_service), cache_service: "BaseCacheService" = Depends(get_cache_service),
user=Depends(get_current_active_user),
): ):
"""Stream the build process based on stored flow data.""" """Stream the build process based on stored flow data."""
@ -167,9 +168,9 @@ async def stream_build(
} }
yield str(StreamData(event="log", data=log_dict)) yield str(StreamData(event="log", data=log_dict))
if vertex.is_task: if vertex.is_task:
vertex = try_running_celery_task(vertex) vertex = try_running_celery_task(vertex, user.id)
else: else:
vertex.build() vertex.build(user_id=user.id)
params = vertex._built_object_repr() params = vertex._built_object_repr()
valid = True valid = True
logger.debug(f"Building node {str(vertex.vertex_type)}") logger.debug(f"Building node {str(vertex.vertex_type)}")
@ -233,7 +234,7 @@ async def stream_build(
raise HTTPException(status_code=500, detail=str(exc)) raise HTTPException(status_code=500, detail=str(exc))
def try_running_celery_task(vertex): def try_running_celery_task(vertex, user_id):
# Try running the task in celery # Try running the task in celery
# and set the task_id to the local vertex # and set the task_id to the local vertex
# if it fails, run the task locally # if it fails, run the task locally
@ -245,5 +246,5 @@ def try_running_celery_task(vertex):
except Exception as exc: except Exception as exc:
logger.debug(f"Error running task in celery: {exc}") logger.debug(f"Error running task in celery: {exc}")
vertex.task_id = None vertex.task_id = None
vertex.build() vertex.build(user_id=user_id)
return vertex return vertex

View file

@ -227,6 +227,7 @@ def get_version():
@router.post("/custom_component", status_code=HTTPStatus.OK) @router.post("/custom_component", status_code=HTTPStatus.OK)
async def custom_component( async def custom_component(
raw_code: CustomComponentCode, raw_code: CustomComponentCode,
user: User = Depends(get_current_active_user),
): ):
from langflow.interface.types import ( from langflow.interface.types import (
build_langchain_template_custom_component, build_langchain_template_custom_component,
@ -235,4 +236,4 @@ async def custom_component(
extractor = CustomComponent(code=raw_code.code) extractor = CustomComponent(code=raw_code.code)
extractor.is_check_valid() extractor.is_check_valid()
return build_langchain_template_custom_component(extractor) return build_langchain_template_custom_component(extractor, user_id=user.id)

View file

@ -1,6 +1,5 @@
import ast import ast
from typing import Any, ClassVar, Optional from typing import Any, ClassVar, Optional
from pydantic import BaseModel
from fastapi import HTTPException from fastapi import HTTPException
from langflow.utils import validate from langflow.utils import validate
@ -15,7 +14,7 @@ class ComponentFunctionEntrypointNameNullError(HTTPException):
pass pass
class Component(BaseModel): class Component:
ERROR_CODE_NULL: ClassVar[str] = "Python code must be provided." ERROR_CODE_NULL: ClassVar[str] = "Python code must be provided."
ERROR_FUNCTION_ENTRYPOINT_NAME_NULL: ClassVar[ ERROR_FUNCTION_ENTRYPOINT_NAME_NULL: ClassVar[
str str
@ -26,7 +25,8 @@ class Component(BaseModel):
field_config: dict = {} field_config: dict = {}
def __init__(self, **data): def __init__(self, **data):
super().__init__(**data) for key, value in data.items():
setattr(self, key, value)
def get_code_tree(self, code: str): def get_code_tree(self, code: str):
parser = CodeParser(code) parser = CodeParser(code)

View file

@ -14,7 +14,7 @@ from langflow.services.database.models.flow import Flow
import yaml import yaml
class CustomComponent(Component, extra="allow"): class CustomComponent(Component):
display_name: Optional[str] = "Custom Component" display_name: Optional[str] = "Custom Component"
description: Optional[str] = "Custom Component" description: Optional[str] = "Custom Component"
code: Optional[str] = None code: Optional[str] = None
@ -203,7 +203,7 @@ class CustomComponent(Component, extra="allow"):
raise ValueError(f"Flow {flow_id} not found") raise ValueError(f"Flow {flow_id} not found")
if tweaks: if tweaks:
graph_data = process_tweaks(graph_data=graph_data, tweaks=tweaks) graph_data = process_tweaks(graph_data=graph_data, tweaks=tweaks)
return build_sorted_vertices(graph_data) return build_sorted_vertices(graph_data, self.user_id)
def list_flows(self, *, get_session: Optional[Callable] = None) -> List[Flow]: def list_flows(self, *, get_session: Optional[Callable] = None) -> List[Flow]:
if not self.user_id: if not self.user_id:

View file

@ -3,7 +3,7 @@ from langflow.graph import Graph
from loguru import logger from loguru import logger
def build_sorted_vertices(data_graph) -> Tuple[Graph, Dict]: def build_sorted_vertices(data_graph, user_id) -> Tuple[Graph, Dict]:
""" """
Build langchain object from data_graph. Build langchain object from data_graph.
""" """
@ -13,7 +13,7 @@ def build_sorted_vertices(data_graph) -> Tuple[Graph, Dict]:
sorted_vertices = graph.topological_sort() sorted_vertices = graph.topological_sort()
artifacts = {} artifacts = {}
for vertex in sorted_vertices: for vertex in sorted_vertices:
vertex.build() vertex.build(user_id=user_id)
if vertex.artifacts: if vertex.artifacts:
artifacts.update(vertex.artifacts) artifacts.update(vertex.artifacts)
return graph, artifacts return graph, artifacts

View file

@ -208,7 +208,7 @@ def update_attributes(frontend_node, template_config):
frontend_node[attribute] = template_config[attribute] frontend_node[attribute] = template_config[attribute]
def build_field_config(custom_component: CustomComponent): def build_field_config(custom_component: CustomComponent, user_id: str = None):
"""Build the field configuration for a custom component""" """Build the field configuration for a custom component"""
try: try:
@ -218,7 +218,7 @@ def build_field_config(custom_component: CustomComponent):
return {} return {}
try: try:
return custom_class().build_config() return custom_class(user_id=user_id).build_config()
except Exception as exc: except Exception as exc:
logger.error(f"Error while building field config: {str(exc)}") logger.error(f"Error while building field config: {str(exc)}")
return {} return {}
@ -306,7 +306,9 @@ def add_output_types(frontend_node, return_types: List[str]):
frontend_node.get("output_types").append(return_type) frontend_node.get("output_types").append(return_type)
def build_langchain_template_custom_component(custom_component: CustomComponent): def build_langchain_template_custom_component(
custom_component: CustomComponent, user_id: str = None
):
"""Build a custom component template for the langchain""" """Build a custom component template for the langchain"""
try: try:
logger.debug("Building custom component template") logger.debug("Building custom component template")
@ -319,7 +321,7 @@ def build_langchain_template_custom_component(custom_component: CustomComponent)
update_attributes(frontend_node, template_config) update_attributes(frontend_node, template_config)
logger.debug("Updated attributes") logger.debug("Updated attributes")
field_config = build_field_config(custom_component) field_config = build_field_config(custom_component, user_id=user_id)
logger.debug("Built field config") logger.debug("Built field config")
entrypoint_args = custom_component.get_function_entrypoint_args entrypoint_args = custom_component.get_function_entrypoint_args