Format code

This commit is contained in:
igorrCarvalho 2024-03-01 16:10:31 -03:00
commit 787bc5111a
19 changed files with 204 additions and 399 deletions

View file

@ -33,9 +33,7 @@ class AsyncStreamingLLMCallbackHandleSIO(AsyncCallbackHandler):
resp = ChatResponse(message=token, type="stream", intermediate_steps="") resp = ChatResponse(message=token, type="stream", intermediate_steps="")
await self.socketio_service.emit_token(to=self.sid, data=resp.model_dump()) await self.socketio_service.emit_token(to=self.sid, data=resp.model_dump())
async def on_tool_start( async def on_tool_start(self, serialized: Dict[str, Any], input_str: str, **kwargs: Any) -> Any:
self, serialized: Dict[str, Any], input_str: str, **kwargs: Any
) -> Any:
"""Run when tool starts running.""" """Run when tool starts running."""
resp = ChatResponse( resp = ChatResponse(
message="", message="",
@ -73,9 +71,7 @@ class AsyncStreamingLLMCallbackHandleSIO(AsyncCallbackHandler):
try: try:
# This is to emulate the stream of tokens # This is to emulate the stream of tokens
for resp in resps: for resp in resps:
await self.socketio_service.emit_token( await self.socketio_service.emit_token(to=self.sid, data=resp.model_dump())
to=self.sid, data=resp.model_dump()
)
except Exception as exc: except Exception as exc:
logger.error(f"Error sending response: {exc}") logger.error(f"Error sending response: {exc}")
@ -101,9 +97,7 @@ class AsyncStreamingLLMCallbackHandleSIO(AsyncCallbackHandler):
resp = PromptResponse( resp = PromptResponse(
prompt=text, prompt=text,
) )
await self.socketio_service.emit_message( await self.socketio_service.emit_message(to=self.sid, data=resp.model_dump())
to=self.sid, data=resp.model_dump()
)
async def on_agent_action(self, action: AgentAction, **kwargs: Any): async def on_agent_action(self, action: AgentAction, **kwargs: Any):
log = f"Thought: {action.log}" log = f"Thought: {action.log}"
@ -113,9 +107,7 @@ class AsyncStreamingLLMCallbackHandleSIO(AsyncCallbackHandler):
logs = log.split("\n") logs = log.split("\n")
for log in logs: for log in logs:
resp = ChatResponse(message="", type="stream", intermediate_steps=log) resp = ChatResponse(message="", type="stream", intermediate_steps=log)
await self.socketio_service.emit_token( await self.socketio_service.emit_token(to=self.sid, data=resp.model_dump())
to=self.sid, data=resp.model_dump()
)
else: else:
resp = ChatResponse(message="", type="stream", intermediate_steps=log) resp = ChatResponse(message="", type="stream", intermediate_steps=log)
await self.socketio_service.emit_token(to=self.sid, data=resp.model_dump()) await self.socketio_service.emit_token(to=self.sid, data=resp.model_dump())

View file

@ -99,12 +99,8 @@ async def build_vertex(
cache = chat_service.get_cache(flow_id) cache = chat_service.get_cache(flow_id)
if not cache: if not cache:
# If there's no cache # If there's no cache
logger.warning( logger.warning(f"No cache found for {flow_id}. Building graph starting at {vertex_id}")
f"No cache found for {flow_id}. Building graph starting at {vertex_id}" graph = build_and_cache_graph(flow_id=flow_id, session=next(get_session()), chat_service=chat_service)
)
graph = build_and_cache_graph(
flow_id=flow_id, session=next(get_session()), chat_service=chat_service
)
else: else:
graph = cache.get("result") graph = cache.get("result")
result_data_response = ResultDataResponse(results={}) result_data_response = ResultDataResponse(results={})
@ -197,9 +193,7 @@ async def build_vertex_stream(
else: else:
graph = cache.get("result") graph = cache.get("result")
else: else:
session_data = await session_service.load_session( session_data = await session_service.load_session(session_id, flow_id=flow_id)
session_id, flow_id=flow_id
)
graph, artifacts = session_data if session_data else (None, None) graph, artifacts = session_data if session_data else (None, None)
if not graph: if not graph:
raise ValueError(f"No graph found for {flow_id}.") raise ValueError(f"No graph found for {flow_id}.")

View file

@ -36,14 +36,8 @@ class LCVectorStoreComponent(CustomComponent):
""" """
docs: List[Document] = [] docs: List[Document] = []
if ( if input_value and isinstance(input_value, str) and hasattr(vector_store, "search"):
input_value docs = vector_store.search(query=input_value, search_type=search_type.lower())
and isinstance(input_value, str)
and hasattr(vector_store, "search")
):
docs = vector_store.search(
query=input_value, search_type=search_type.lower()
)
else: else:
raise ValueError("Invalid inputs provided.") raise ValueError("Invalid inputs provided.")
return docs_to_records(docs) return docs_to_records(docs)

View file

@ -86,9 +86,7 @@ class Graph:
if getattr(vertex, attribute): if getattr(vertex, attribute):
getattr(self, f"_{attribute}_vertices").append(vertex.id) getattr(self, f"_{attribute}_vertices").append(vertex.id)
async def _run( async def _run(self, inputs: Dict[str, str], stream: bool) -> List[Optional["ResultData"]]:
self, inputs: Dict[str, str], stream: bool
) -> List[Optional["ResultData"]]:
"""Runs the graph with the given inputs.""" """Runs the graph with the given inputs."""
for vertex_id in self._is_input_vertices: for vertex_id in self._is_input_vertices:
vertex = self.get_vertex(vertex_id) vertex = self.get_vertex(vertex_id)
@ -111,9 +109,7 @@ class Graph:
outputs.append(vertex.result) outputs.append(vertex.result)
return outputs return outputs
async def run( async def run(self, inputs: Dict[str, Union[str, list[str]]], stream: bool) -> List[Optional["ResultData"]]:
self, inputs: Dict[str, Union[str, list[str]]], stream: bool
) -> List[Optional["ResultData"]]:
"""Runs the graph with the given inputs.""" """Runs the graph with the given inputs."""
# inputs is {"message": "Hello, world!"} # inputs is {"message": "Hello, world!"}
@ -125,9 +121,7 @@ class Graph:
if not isinstance(inputs_values, list): if not isinstance(inputs_values, list):
inputs_values = [inputs_values] inputs_values = [inputs_values]
for input_value in inputs_values: for input_value in inputs_values:
run_outputs = await self._run( run_outputs = await self._run({INPUT_FIELD_NAME: input_value}, stream=stream)
{INPUT_FIELD_NAME: input_value}, stream=stream
)
logger.debug(f"Run outputs: {run_outputs}") logger.debug(f"Run outputs: {run_outputs}")
outputs.extend(run_outputs) outputs.extend(run_outputs)
return outputs return outputs
@ -167,9 +161,7 @@ class Graph:
def build_parent_child_map(self): def build_parent_child_map(self):
parent_child_map = defaultdict(list) parent_child_map = defaultdict(list)
for vertex in self.vertices: for vertex in self.vertices:
parent_child_map[vertex.id] = [ parent_child_map[vertex.id] = [child.id for child in self.get_successors(vertex)]
child.id for child in self.get_successors(vertex)
]
return parent_child_map return parent_child_map
def increment_run_count(self): def increment_run_count(self):
@ -344,11 +336,7 @@ class Graph:
return return
self.vertices.remove(vertex) self.vertices.remove(vertex)
self.vertex_map.pop(vertex_id) self.vertex_map.pop(vertex_id)
self.edges = [ self.edges = [edge for edge in self.edges if edge.source_id != vertex_id and edge.target_id != vertex_id]
edge
for edge in self.edges
if edge.source_id != vertex_id and edge.target_id != vertex_id
]
def _build_vertex_params(self) -> None: def _build_vertex_params(self) -> None:
"""Identifies and handles the LLM vertex within the graph.""" """Identifies and handles the LLM vertex within the graph."""
@ -369,9 +357,7 @@ class Graph:
return return
for vertex in self.vertices: for vertex in self.vertices:
if not self._validate_vertex(vertex): if not self._validate_vertex(vertex):
raise ValueError( raise ValueError(f"{vertex.display_name} is not connected to any other components")
f"{vertex.display_name} is not connected to any other components"
)
def _validate_vertex(self, vertex: Vertex) -> bool: def _validate_vertex(self, vertex: Vertex) -> bool:
"""Validates a vertex.""" """Validates a vertex."""
@ -428,9 +414,7 @@ class Graph:
tasks = [] tasks = []
for vertex_id in layer: for vertex_id in layer:
vertex = self.get_vertex(vertex_id) vertex = self.get_vertex(vertex_id)
task = asyncio.create_task( task = asyncio.create_task(vertex.build(), name=f"layer-{layer_index}-vertex-{vertex_id}")
vertex.build(), name=f"layer-{layer_index}-vertex-{vertex_id}"
)
tasks.append(task) tasks.append(task)
logger.debug(f"Running layer {layer_index} with {len(tasks)} tasks") logger.debug(f"Running layer {layer_index} with {len(tasks)} tasks")
await self._execute_tasks(tasks) await self._execute_tasks(tasks)
@ -469,9 +453,7 @@ class Graph:
def dfs(vertex): def dfs(vertex):
if state[vertex] == 1: if state[vertex] == 1:
# We have a cycle # We have a cycle
raise ValueError( raise ValueError("Graph contains a cycle, cannot perform topological sort")
"Graph contains a cycle, cannot perform topological sort"
)
if state[vertex] == 0: if state[vertex] == 0:
state[vertex] = 1 state[vertex] = 1
for edge in vertex.edges: for edge in vertex.edges:
@ -495,17 +477,11 @@ class Graph:
def get_predecessors(self, vertex): def get_predecessors(self, vertex):
"""Returns the predecessors of a vertex.""" """Returns the predecessors of a vertex."""
return [ return [self.get_vertex(source_id) for source_id in self.predecessor_map.get(vertex.id, [])]
self.get_vertex(source_id)
for source_id in self.predecessor_map.get(vertex.id, [])
]
def get_successors(self, vertex): def get_successors(self, vertex):
"""Returns the successors of a vertex.""" """Returns the successors of a vertex."""
return [ return [self.get_vertex(target_id) for target_id in self.successor_map.get(vertex.id, [])]
self.get_vertex(target_id)
for target_id in self.successor_map.get(vertex.id, [])
]
def get_vertex_neighbors(self, vertex: Vertex) -> Dict[Vertex, int]: def get_vertex_neighbors(self, vertex: Vertex) -> Dict[Vertex, int]:
"""Returns the neighbors of a vertex.""" """Returns the neighbors of a vertex."""
@ -544,9 +520,7 @@ class Graph:
edges.append(ContractEdge(source, target, edge)) edges.append(ContractEdge(source, target, edge))
return edges return edges
def _get_vertex_class( def _get_vertex_class(self, node_type: str, node_base_type: str, node_id: str) -> Type[Vertex]:
self, node_type: str, node_base_type: str, node_id: str
) -> Type[Vertex]:
"""Returns the node class based on the node type.""" """Returns the node class based on the node type."""
# First we check for the node_base_type # First we check for the node_base_type
node_name = node_id.split("-")[0] node_name = node_id.split("-")[0]
@ -577,18 +551,14 @@ class Graph:
vertex_type: str = vertex_data["type"] # type: ignore vertex_type: str = vertex_data["type"] # type: ignore
vertex_base_type: str = vertex_data["node"]["template"]["_type"] # type: ignore vertex_base_type: str = vertex_data["node"]["template"]["_type"] # type: ignore
VertexClass = self._get_vertex_class( VertexClass = self._get_vertex_class(vertex_type, vertex_base_type, vertex_data["id"])
vertex_type, vertex_base_type, vertex_data["id"]
)
vertex_instance = VertexClass(vertex, graph=self) vertex_instance = VertexClass(vertex, graph=self)
vertex_instance.set_top_level(self.top_level_vertices) vertex_instance.set_top_level(self.top_level_vertices)
vertices.append(vertex_instance) vertices.append(vertex_instance)
return vertices return vertices
def get_children_by_vertex_type( def get_children_by_vertex_type(self, vertex: Vertex, vertex_type: str) -> List[Vertex]:
self, vertex: Vertex, vertex_type: str
) -> List[Vertex]:
"""Returns the children of a vertex based on the vertex type.""" """Returns the children of a vertex based on the vertex type."""
children = [] children = []
vertex_types = [vertex.data["type"]] vertex_types = [vertex.data["type"]]
@ -600,9 +570,7 @@ class Graph:
def __repr__(self): def __repr__(self):
vertex_ids = [vertex.id for vertex in self.vertices] vertex_ids = [vertex.id for vertex in self.vertices]
edges_repr = "\n".join( edges_repr = "\n".join([f"{edge.source_id} --> {edge.target_id}" for edge in self.edges])
[f"{edge.source_id} --> {edge.target_id}" for edge in self.edges]
)
return f"Graph:\nNodes: {vertex_ids}\nConnections:\n{edges_repr}" return f"Graph:\nNodes: {vertex_ids}\nConnections:\n{edges_repr}"
def sort_up_to_vertex(self, vertex_id: str) -> List[Vertex]: def sort_up_to_vertex(self, vertex_id: str) -> List[Vertex]:
@ -633,9 +601,7 @@ class Graph:
"""Performs a layered topological sort of the vertices in the graph.""" """Performs a layered topological sort of the vertices in the graph."""
vertices_ids = {vertex.id for vertex in vertices} vertices_ids = {vertex.id for vertex in vertices}
# Queue for vertices with no incoming edges # Queue for vertices with no incoming edges
queue = deque( queue = deque(vertex.id for vertex in vertices if self.in_degree_map[vertex.id] == 0)
vertex.id for vertex in vertices if self.in_degree_map[vertex.id] == 0
)
layers: List[List[str]] = [] layers: List[List[str]] = []
current_layer = 0 current_layer = 0
@ -698,9 +664,7 @@ class Graph:
return refined_layers return refined_layers
def sort_chat_inputs_first( def sort_chat_inputs_first(self, vertices_layers: List[List[str]]) -> List[List[str]]:
self, vertices_layers: List[List[str]]
) -> List[List[str]]:
chat_inputs_first = [] chat_inputs_first = []
for layer in vertices_layers: for layer in vertices_layers:
for vertex_id in layer: for vertex_id in layer:
@ -729,15 +693,11 @@ class Graph:
self._sorted_vertices_layers = vertices_layers self._sorted_vertices_layers = vertices_layers
return vertices_layers return vertices_layers
def sort_interface_components_first( def sort_interface_components_first(self, vertices_layers: List[List[str]]) -> List[List[str]]:
self, vertices_layers: List[List[str]]
) -> List[List[str]]:
"""Sorts the vertices in the graph so that vertices containing ChatInput or ChatOutput come first.""" """Sorts the vertices in the graph so that vertices containing ChatInput or ChatOutput come first."""
def contains_interface_component(vertex): def contains_interface_component(vertex):
return any( return any(component.value in vertex for component in InterfaceComponentTypes)
component.value in vertex for component in InterfaceComponentTypes
)
# Sort each inner list so that vertices containing ChatInput or ChatOutput come first # Sort each inner list so that vertices containing ChatInput or ChatOutput come first
sorted_vertices = [ sorted_vertices = [
@ -749,22 +709,16 @@ class Graph:
] ]
return sorted_vertices return sorted_vertices
def sort_by_avg_build_time( def sort_by_avg_build_time(self, vertices_layers: List[List[str]]) -> List[List[str]]:
self, vertices_layers: List[List[str]]
) -> List[List[str]]:
"""Sorts the vertices in the graph so that vertices with the lowest average build time come first.""" """Sorts the vertices in the graph so that vertices with the lowest average build time come first."""
def sort_layer_by_avg_build_time(vertices_ids: List[str]) -> List[str]: def sort_layer_by_avg_build_time(vertices_ids: List[str]) -> List[str]:
"""Sorts the vertices in the graph so that vertices with the lowest average build time come first.""" """Sorts the vertices in the graph so that vertices with the lowest average build time come first."""
if len(vertices_ids) == 1: if len(vertices_ids) == 1:
return vertices_ids return vertices_ids
vertices_ids.sort( vertices_ids.sort(key=lambda vertex_id: self.get_vertex(vertex_id).avg_build_time)
key=lambda vertex_id: self.get_vertex(vertex_id).avg_build_time
)
return vertices_ids return vertices_ids
sorted_vertices = [ sorted_vertices = [sort_layer_by_avg_build_time(layer) for layer in vertices_layers]
sort_layer_by_avg_build_time(layer) for layer in vertices_layers
]
return sorted_vertices return sorted_vertices

View file

@ -47,13 +47,8 @@ class Vertex:
self.will_stream = False self.will_stream = False
self.updated_raw_params = False self.updated_raw_params = False
self.id: str = data["id"] self.id: str = data["id"]
self.is_input = any( self.is_input = any(input_component_name in self.id for input_component_name in INPUT_COMPONENTS)
input_component_name in self.id for input_component_name in INPUT_COMPONENTS self.is_output = any(output_component_name in self.id for output_component_name in OUTPUT_COMPONENTS)
)
self.is_output = any(
output_component_name in self.id
for output_component_name in OUTPUT_COMPONENTS
)
self.has_session_id = None self.has_session_id = None
self._custom_component = None self._custom_component = None
self.has_external_input = False self.has_external_input = False
@ -97,17 +92,11 @@ class Vertex:
def set_state(self, state: str): def set_state(self, state: str):
self.state = VertexStates[state] self.state = VertexStates[state]
if ( if self.state == VertexStates.INACTIVE and self.graph.in_degree_map[self.id] < 2:
self.state == VertexStates.INACTIVE
and self.graph.in_degree_map[self.id] < 2
):
# If the vertex is inactive and has only one in degree # If the vertex is inactive and has only one in degree
# it means that it is not a merge point in the graph # it means that it is not a merge point in the graph
self.graph.inactive_vertices.add(self.id) self.graph.inactive_vertices.add(self.id)
elif ( elif self.state == VertexStates.ACTIVE and self.id in self.graph.inactive_vertices:
self.state == VertexStates.ACTIVE
and self.id in self.graph.inactive_vertices
):
self.graph.inactive_vertices.remove(self.id) self.graph.inactive_vertices.remove(self.id)
@property @property
@ -124,9 +113,7 @@ class Vertex:
# If the Vertex.type is a power component # If the Vertex.type is a power component
# then we need to return the built object # then we need to return the built object
# instead of the result dict # instead of the result dict
if self.is_interface_component and not isinstance( if self.is_interface_component and not isinstance(self._built_object, UnbuiltObject):
self._built_object, UnbuiltObject
):
result = self._built_object result = self._built_object
# if it is not a dict or a string and hasattr model_dump then # if it is not a dict or a string and hasattr model_dump then
# return the model_dump # return the model_dump
@ -136,11 +123,7 @@ class Vertex:
if isinstance(self._built_result, UnbuiltResult): if isinstance(self._built_result, UnbuiltResult):
return {} return {}
return ( return self._built_result if isinstance(self._built_result, dict) else {"result": self._built_result}
self._built_result
if isinstance(self._built_result, dict)
else {"result": self._built_result}
)
def set_artifacts(self) -> None: def set_artifacts(self) -> None:
pass pass
@ -206,31 +189,19 @@ class Vertex:
self.selected_output_type = self.data["node"].get("selected_output_type") self.selected_output_type = self.data["node"].get("selected_output_type")
self.is_input = self.data["node"].get("is_input") or self.is_input self.is_input = self.data["node"].get("is_input") or self.is_input
self.is_output = self.data["node"].get("is_output") or self.is_output self.is_output = self.data["node"].get("is_output") or self.is_output
template_dicts = { template_dicts = {key: value for key, value in self.data["node"]["template"].items() if isinstance(value, dict)}
key: value
for key, value in self.data["node"]["template"].items()
if isinstance(value, dict)
}
self.has_session_id = "session_id" in template_dicts self.has_session_id = "session_id" in template_dicts
self.required_inputs = [ self.required_inputs = [
template_dicts[key]["type"] template_dicts[key]["type"] for key, value in template_dicts.items() if value["required"]
for key, value in template_dicts.items()
if value["required"]
] ]
self.optional_inputs = [ self.optional_inputs = [
template_dicts[key]["type"] template_dicts[key]["type"] for key, value in template_dicts.items() if not value["required"]
for key, value in template_dicts.items()
if not value["required"]
] ]
# Add the template_dicts[key]["input_types"] to the optional_inputs # Add the template_dicts[key]["input_types"] to the optional_inputs
self.optional_inputs.extend( self.optional_inputs.extend(
[ [input_type for value in template_dicts.values() for input_type in value.get("input_types", [])]
input_type
for value in template_dicts.values()
for input_type in value.get("input_types", [])
]
) )
template_dict = self.data["node"]["template"] template_dict = self.data["node"]["template"]
@ -277,11 +248,7 @@ class Vertex:
self.updated_raw_params = False self.updated_raw_params = False
return return
template_dict = { template_dict = {key: value for key, value in self.data["node"]["template"].items() if isinstance(value, dict)}
key: value
for key, value in self.data["node"]["template"].items()
if isinstance(value, dict)
}
params = {} params = {}
for edge in self.edges: for edge in self.edges:
@ -332,11 +299,7 @@ class Vertex:
# list of dicts, so we need to convert it to a dict # list of dicts, so we need to convert it to a dict
# before passing it to the build method # before passing it to the build method
if isinstance(val, list): if isinstance(val, list):
params[key] = { params[key] = {k: v for item in value.get("value", []) for k, v in item.items()}
k: v
for item in value.get("value", [])
for k, v in item.items()
}
elif isinstance(val, dict): elif isinstance(val, dict):
params[key] = val params[key] = val
elif value.get("type") == "int" and val is not None: elif value.get("type") == "int" and val is not None:
@ -429,9 +392,7 @@ class Vertex:
if isinstance(self._built_object, str): if isinstance(self._built_object, str):
self._built_result = self._built_object self._built_result = self._built_object
result = await generate_result( result = await generate_result(self._built_object, inputs, self.has_external_output, session_id)
self._built_object, inputs, self.has_external_output, session_id
)
self._built_result = result self._built_result = result
async def _build_each_node_in_params_dict(self, user_id=None): async def _build_each_node_in_params_dict(self, user_id=None):
@ -461,9 +422,7 @@ class Vertex:
""" """
return all(self._is_node(node) for node in value) return all(self._is_node(node) for node in value)
async def get_result( async def get_result(self, requester: Optional["Vertex"] = None, user_id=None, timeout=None) -> Any:
self, requester: Optional["Vertex"] = None, user_id=None, timeout=None
) -> Any:
# PLEASE REVIEW THIS IF STATEMENT # PLEASE REVIEW THIS IF STATEMENT
# Check if the Vertex was built already # Check if the Vertex was built already
if self._built: if self._built:
@ -497,9 +456,7 @@ class Vertex:
self._extend_params_list_with_result(key, result) self._extend_params_list_with_result(key, result)
self.params[key] = result self.params[key] = result
async def _build_list_of_nodes_and_update_params( async def _build_list_of_nodes_and_update_params(self, key, nodes: List["Vertex"], user_id=None):
self, key, nodes: List["Vertex"], user_id=None
):
""" """
Iterates over a list of nodes, builds each and updates the params dictionary. Iterates over a list of nodes, builds each and updates the params dictionary.
""" """
@ -553,9 +510,7 @@ class Vertex:
except Exception as exc: except Exception as exc:
logger.exception(exc) logger.exception(exc)
raise ValueError( raise ValueError(f"Error building node {self.display_name}: {str(exc)}") from exc
f"Error building node {self.display_name}: {str(exc)}"
) from exc
def _update_built_object_and_artifacts(self, result): def _update_built_object_and_artifacts(self, result):
""" """
@ -635,24 +590,16 @@ class Vertex:
return self._built_object return self._built_object
# Get the requester edge # Get the requester edge
requester_edge = next( requester_edge = next((edge for edge in self.edges if edge.target_id == requester.id), None)
(edge for edge in self.edges if edge.target_id == requester.id), None
)
# Return the result of the requester edge # Return the result of the requester edge
return ( return None if requester_edge is None else await requester_edge.get_result(source=self, target=requester)
None
if requester_edge is None
else await requester_edge.get_result(source=self, target=requester)
)
def add_edge(self, edge: "ContractEdge") -> None: def add_edge(self, edge: "ContractEdge") -> None:
if edge not in self.edges: if edge not in self.edges:
self.edges.append(edge) self.edges.append(edge)
def __repr__(self) -> str: def __repr__(self) -> str:
return ( return f"Vertex(display_name={self.display_name}, id={self.id}, data={self.data})"
f"Vertex(display_name={self.display_name}, id={self.id}, data={self.data})"
)
def __eq__(self, __o: object) -> bool: def __eq__(self, __o: object) -> bool:
try: try:
@ -665,11 +612,7 @@ class Vertex:
def _built_object_repr(self): def _built_object_repr(self):
# Add a message with an emoji, stars for sucess, # Add a message with an emoji, stars for sucess,
return ( return "Built sucessfully ✨" if self._built_object is not None else "Failed to build 😵‍💫"
"Built sucessfully ✨"
if self._built_object is not None
else "Failed to build 😵‍💫"
)
class StatefulVertex(Vertex): class StatefulVertex(Vertex):

View file

@ -27,18 +27,14 @@ from langflow.utils import validate
from langflow.utils.util import get_base_classes from langflow.utils.util import get_base_classes
def add_output_types( def add_output_types(frontend_node: CustomComponentFrontendNode, return_types: List[str]):
frontend_node: CustomComponentFrontendNode, return_types: List[str]
):
"""Add output types to the frontend node""" """Add output types to the frontend node"""
for return_type in return_types: for return_type in return_types:
if return_type is None: if return_type is None:
raise HTTPException( raise HTTPException(
status_code=400, status_code=400,
detail={ detail={
"error": ( "error": ("Invalid return type. Please check your code and try again."),
"Invalid return type. Please check your code and try again."
),
"traceback": traceback.format_exc(), "traceback": traceback.format_exc(),
}, },
) )
@ -69,18 +65,14 @@ def reorder_fields(frontend_node: CustomComponentFrontendNode, field_order: List
frontend_node.template.fields = reordered_fields frontend_node.template.fields = reordered_fields
def add_base_classes( def add_base_classes(frontend_node: CustomComponentFrontendNode, return_types: List[str]):
frontend_node: CustomComponentFrontendNode, return_types: List[str]
):
"""Add base classes to the frontend node""" """Add base classes to the frontend node"""
for return_type_instance in return_types: for return_type_instance in return_types:
if return_type_instance is None: if return_type_instance is None:
raise HTTPException( raise HTTPException(
status_code=400, status_code=400,
detail={ detail={
"error": ( "error": ("Invalid return type. Please check your code and try again."),
"Invalid return type. Please check your code and try again."
),
"traceback": traceback.format_exc(), "traceback": traceback.format_exc(),
}, },
) )
@ -155,14 +147,10 @@ def add_new_custom_field(
# If options is a list, then it's a dropdown # If options is a list, then it's a dropdown
# If options is None, then it's a list of strings # If options is None, then it's a list of strings
is_list = isinstance(field_config.get("options"), list) is_list = isinstance(field_config.get("options"), list)
field_config["is_list"] = ( field_config["is_list"] = is_list or field_config.get("is_list", False) or field_contains_list
is_list or field_config.get("is_list", False) or field_contains_list
)
if "name" in field_config: if "name" in field_config:
warnings.warn( warnings.warn("The 'name' key in field_config is used to build the object and can't be changed.")
"The 'name' key in field_config is used to build the object and can't be changed."
)
required = field_config.pop("required", field_required) required = field_config.pop("required", field_required)
placeholder = field_config.pop("placeholder", "") placeholder = field_config.pop("placeholder", "")
@ -193,9 +181,7 @@ def add_extra_fields(frontend_node, field_config, function_args):
if "name" not in extra_field or extra_field["name"] == "self": if "name" not in extra_field or extra_field["name"] == "self":
continue continue
field_name, field_type, field_value, field_required = get_field_properties( field_name, field_type, field_value, field_required = get_field_properties(extra_field)
extra_field
)
config = field_config.get(field_name, {}) config = field_config.get(field_name, {})
frontend_node = add_new_custom_field( frontend_node = add_new_custom_field(
frontend_node, frontend_node,
@ -233,9 +219,7 @@ def run_build_config(
raise HTTPException( raise HTTPException(
status_code=400, status_code=400,
detail={ detail={
"error": ( "error": ("Invalid type convertion. Please check your code and try again."),
"Invalid type convertion. Please check your code and try again."
),
"traceback": traceback.format_exc(), "traceback": traceback.format_exc(),
}, },
) from exc ) from exc
@ -263,9 +247,7 @@ def run_build_config(
raise HTTPException( raise HTTPException(
status_code=400, status_code=400,
detail={ detail={
"error": ( "error": ("Invalid type convertion. Please check your code and try again."),
"Invalid type convertion. Please check your code and try again."
),
"traceback": traceback.format_exc(), "traceback": traceback.format_exc(),
}, },
) from exc ) from exc
@ -318,24 +300,16 @@ def build_custom_component_template(
try: try:
frontend_node = build_frontend_node(custom_component.template_config) frontend_node = build_frontend_node(custom_component.template_config)
field_config, custom_instance = run_build_config( field_config, custom_instance = run_build_config(custom_component, user_id=user_id, update_field=update_field)
custom_component, user_id=user_id, update_field=update_field
)
entrypoint_args = custom_component.get_function_entrypoint_args entrypoint_args = custom_component.get_function_entrypoint_args
add_extra_fields(frontend_node, field_config, entrypoint_args) add_extra_fields(frontend_node, field_config, entrypoint_args)
frontend_node = add_code_field( frontend_node = add_code_field(frontend_node, custom_component.code, field_config.get("code", {}))
frontend_node, custom_component.code, field_config.get("code", {})
)
add_base_classes( add_base_classes(frontend_node, custom_component.get_function_entrypoint_return_type)
frontend_node, custom_component.get_function_entrypoint_return_type add_output_types(frontend_node, custom_component.get_function_entrypoint_return_type)
)
add_output_types(
frontend_node, custom_component.get_function_entrypoint_return_type
)
reorder_fields(frontend_node, custom_instance._get_field_order()) reorder_fields(frontend_node, custom_instance._get_field_order())
@ -346,9 +320,7 @@ def build_custom_component_template(
raise HTTPException( raise HTTPException(
status_code=400, status_code=400,
detail={ detail={
"error": ( "error": ("Invalid type convertion. Please check your code and try again."),
"Invalid type convertion. Please check your code and try again."
),
"traceback": traceback.format_exc(), "traceback": traceback.format_exc(),
}, },
) from exc ) from exc
@ -372,9 +344,7 @@ def build_custom_components(settings_service):
if not settings_service.settings.COMPONENTS_PATH: if not settings_service.settings.COMPONENTS_PATH:
return {} return {}
logger.info( logger.info(f"Building custom components from {settings_service.settings.COMPONENTS_PATH}")
f"Building custom components from {settings_service.settings.COMPONENTS_PATH}"
)
custom_components_from_file = {} custom_components_from_file = {}
processed_paths = set() processed_paths = set()
for path in settings_service.settings.COMPONENTS_PATH: for path in settings_service.settings.COMPONENTS_PATH:
@ -385,9 +355,7 @@ def build_custom_components(settings_service):
custom_component_dict = build_custom_component_list_from_path(path_str) custom_component_dict = build_custom_component_list_from_path(path_str)
if custom_component_dict: if custom_component_dict:
category = next(iter(custom_component_dict)) category = next(iter(custom_component_dict))
logger.info( logger.info(f"Loading {len(custom_component_dict[category])} component(s) from category {category}")
f"Loading {len(custom_component_dict[category])} component(s) from category {category}"
)
custom_components_from_file = merge_nested_dicts_with_renaming( custom_components_from_file = merge_nested_dicts_with_renaming(
custom_components_from_file, custom_component_dict custom_components_from_file, custom_component_dict
) )

View file

@ -22,9 +22,7 @@ async def process_graph(
if build_result is None: if build_result is None:
# Raise user facing error # Raise user facing error
raise ValueError( raise ValueError("There was an error loading the langchain_object. Please, check all the nodes and try again.")
"There was an error loading the langchain_object. Please, check all the nodes and try again."
)
# Generate result and thought # Generate result and thought
try: try:
@ -50,7 +48,5 @@ async def process_graph(
raise e raise e
async def run_build_result( async def run_build_result(build_result: Any, chat_inputs: ChatMessage, client_id: str, session_id: str):
build_result: Any, chat_inputs: ChatMessage, client_id: str, session_id: str
):
return build_result(inputs=chat_inputs.message) return build_result(inputs=chat_inputs.message)

View file

@ -92,16 +92,12 @@ def get_or_create_super_user(session: Session, username, password, is_default):
) )
return None return None
else: else:
logger.debug( logger.debug("User with superuser credentials exists but is not a superuser.")
"User with superuser credentials exists but is not a superuser."
)
return None return None
if user: if user:
if verify_password(password, user.password): if verify_password(password, user.password):
raise ValueError( raise ValueError("User with superuser credentials exists but is not a superuser.")
"User with superuser credentials exists but is not a superuser."
)
else: else:
raise ValueError("Incorrect superuser credentials") raise ValueError("Incorrect superuser credentials")
@ -130,21 +126,15 @@ def setup_superuser(settings_service, session: Session):
username = settings_service.auth_settings.SUPERUSER username = settings_service.auth_settings.SUPERUSER
password = settings_service.auth_settings.SUPERUSER_PASSWORD password = settings_service.auth_settings.SUPERUSER_PASSWORD
is_default = (username == DEFAULT_SUPERUSER) and ( is_default = (username == DEFAULT_SUPERUSER) and (password == DEFAULT_SUPERUSER_PASSWORD)
password == DEFAULT_SUPERUSER_PASSWORD
)
try: try:
user = get_or_create_super_user( user = get_or_create_super_user(session=session, username=username, password=password, is_default=is_default)
session=session, username=username, password=password, is_default=is_default
)
if user is not None: if user is not None:
logger.debug("Superuser created successfully.") logger.debug("Superuser created successfully.")
except Exception as exc: except Exception as exc:
logger.exception(exc) logger.exception(exc)
raise RuntimeError( raise RuntimeError("Could not create superuser. Please create a superuser manually.") from exc
"Could not create superuser. Please create a superuser manually."
) from exc
finally: finally:
settings_service.auth_settings.reset_credentials() settings_service.auth_settings.reset_credentials()
@ -158,9 +148,7 @@ def teardown_superuser(settings_service, session):
if not settings_service.auth_settings.AUTO_LOGIN: if not settings_service.auth_settings.AUTO_LOGIN:
try: try:
logger.debug( logger.debug("AUTO_LOGIN is set to False. Removing default superuser if exists.")
"AUTO_LOGIN is set to False. Removing default superuser if exists."
)
username = DEFAULT_SUPERUSER username = DEFAULT_SUPERUSER
from langflow.services.database.models.user.model import User from langflow.services.database.models.user.model import User
@ -210,9 +198,7 @@ def initialize_session_service():
initialize_settings_service() initialize_settings_service()
service_manager.register_factory( service_manager.register_factory(cache_factory.CacheServiceFactory(), dependencies=[ServiceType.SETTINGS_SERVICE])
cache_factory.CacheServiceFactory(), dependencies=[ServiceType.SETTINGS_SERVICE]
)
service_manager.register_factory( service_manager.register_factory(
session_service_factory.SessionServiceFactory(), session_service_factory.SessionServiceFactory(),
@ -229,9 +215,7 @@ def initialize_services(fix_migration: bool = False, socketio_server=None):
service_manager.register_factory(factory, dependencies=dependencies) service_manager.register_factory(factory, dependencies=dependencies)
except Exception as exc: except Exception as exc:
logger.exception(exc) logger.exception(exc)
raise RuntimeError( raise RuntimeError("Could not initialize services. Please check your settings.") from exc
"Could not initialize services. Please check your settings."
) from exc
# Test cache connection # Test cache connection
service_manager.get(ServiceType.CACHE_SERVICE) service_manager.get(ServiceType.CACHE_SERVICE)
@ -241,9 +225,7 @@ def initialize_services(fix_migration: bool = False, socketio_server=None):
except Exception as exc: except Exception as exc:
logger.error(exc) logger.error(exc)
raise exc raise exc
setup_superuser( setup_superuser(service_manager.get(ServiceType.SETTINGS_SERVICE), next(get_session()))
service_manager.get(ServiceType.SETTINGS_SERVICE), next(get_session())
)
try: try:
get_db_service().migrate_flows_if_auto_login() get_db_service().migrate_flows_if_auto_login()
except Exception as exc: except Exception as exc:

View file

@ -18,11 +18,9 @@ import { getHealth } from "./controllers/API";
import Router from "./routes"; import Router from "./routes";
import useAlertStore from "./stores/alertStore"; import useAlertStore from "./stores/alertStore";
import { useDarkStore } from "./stores/darkStore"; import { useDarkStore } from "./stores/darkStore";
import useFlowStore from "./stores/flowStore";
import useFlowsManagerStore from "./stores/flowsManagerStore"; import useFlowsManagerStore from "./stores/flowsManagerStore";
import { useStoreStore } from "./stores/storeStore"; import { useStoreStore } from "./stores/storeStore";
import { useTypesStore } from "./stores/typesStore"; import { useTypesStore } from "./stores/typesStore";
import { useShortcutsStore } from "./stores/shortcuts";
export default function App() { export default function App() {
const removeFromTempNotificationList = useAlertStore( const removeFromTempNotificationList = useAlertStore(

View file

@ -26,7 +26,7 @@ export default function CrashErrorComponent({
<CardContent className="grid"> <CardContent className="grid">
<div> <div>
<p> <p>
Please report errors with detailed tracebacks on the{" "} Please report errors with detailed tracebacks on the{" "}
<a <a
href="https://github.com/logspace-ai/langflow/issues" href="https://github.com/logspace-ai/langflow/issues"
target="_blank" target="_blank"
@ -37,7 +37,7 @@ export default function CrashErrorComponent({
</a>{" "} </a>{" "}
page. page.
<br></br> <br></br>
Thank you! Thank you!
</p> </p>
</div> </div>
</CardContent> </CardContent>

View file

@ -1,4 +1,4 @@
import { useEffect, useState } from "react"; import { useState } from "react";
import { import {
DropdownMenu, DropdownMenu,
DropdownMenuContent, DropdownMenuContent,
@ -8,7 +8,7 @@ import {
} from "../../../ui/dropdown-menu"; } from "../../../ui/dropdown-menu";
import { useNavigate } from "react-router-dom"; import { useNavigate } from "react-router-dom";
import { Node, useEdges } from "reactflow"; import { Node } from "reactflow";
import { savedHover } from "../../../../constants/constants"; import { savedHover } from "../../../../constants/constants";
import FlowSettingsModal from "../../../../modals/flowSettingsModal"; import FlowSettingsModal from "../../../../modals/flowSettingsModal";
import useAlertStore from "../../../../stores/alertStore"; import useAlertStore from "../../../../stores/alertStore";
@ -110,10 +110,10 @@ export const MenuBar = ({
<IconComponent name="Undo" className="header-menu-options " /> <IconComponent name="Undo" className="header-menu-options " />
Undo Undo
<IconComponent <IconComponent
name="Command" name="Command"
className="absolute right-[1.15rem] top-[0.65em] h-3.5 w-3.5 stroke-2" className="absolute right-[1.15rem] top-[0.65em] h-3.5 w-3.5 stroke-2"
></IconComponent> ></IconComponent>
<span className="absolute right-2 top-[0.4em]">Z</span> <span className="absolute right-2 top-[0.4em]">Z</span>
</DropdownMenuItem> </DropdownMenuItem>
<DropdownMenuItem <DropdownMenuItem
onClick={() => { onClick={() => {
@ -124,10 +124,10 @@ export const MenuBar = ({
<IconComponent name="Redo" className="header-menu-options " /> <IconComponent name="Redo" className="header-menu-options " />
Redo Redo
<IconComponent <IconComponent
name="Command" name="Command"
className="absolute right-[1.15rem] top-[0.65em] h-3.5 w-3.5 stroke-2" className="absolute right-[1.15rem] top-[0.65em] h-3.5 w-3.5 stroke-2"
></IconComponent> ></IconComponent>
<span className="absolute right-2 top-[0.4em]">Y</span> <span className="absolute right-2 top-[0.4em]">Y</span>
</DropdownMenuItem> </DropdownMenuItem>
</DropdownMenuContent> </DropdownMenuContent>
</DropdownMenu> </DropdownMenu>

View file

@ -50,7 +50,6 @@ export default function CodeAreaModal({
} | null>(null); } | null>(null);
const [open, setOpen] = useState(false); const [open, setOpen] = useState(false);
const nodes = useFlowStore((state) => state.nodes); const nodes = useFlowStore((state) => state.nodes);
useEffect(() => { useEffect(() => {
// if nodeClass.template has more fields other than code and dynamic is true // if nodeClass.template has more fields other than code and dynamic is true
@ -147,7 +146,6 @@ export default function CodeAreaModal({
}; };
}, [error, setHeight]); }, [error, setHeight]);
useEffect(() => { useEffect(() => {
setCode(value); setCode(value);
}, [value, open]); }, [value, open]);

View file

@ -488,4 +488,4 @@ export default function Page({
</main> </main>
</div> </div>
); );
} }

View file

@ -24,7 +24,6 @@ import {
createFlowComponent, createFlowComponent,
downloadNode, downloadNode,
expandGroupNode, expandGroupNode,
unselectAllNodes,
updateFlowPosition, updateFlowPosition,
} from "../../../../utils/reactflowUtils"; } from "../../../../utils/reactflowUtils";
import { classNames, cn } from "../../../../utils/utils"; import { classNames, cn } from "../../../../utils/utils";
@ -69,7 +68,7 @@ export default function NodeToolbarComponent({
const edges = useFlowStore((state) => state.edges); const edges = useFlowStore((state) => state.edges);
const setNodes = useFlowStore((state) => state.setNodes); const setNodes = useFlowStore((state) => state.setNodes);
const setEdges = useFlowStore((state) => state.setEdges); const setEdges = useFlowStore((state) => state.setEdges);
const unselectAll = useFlowStore(state => state.unselectAll); const unselectAll = useFlowStore((state) => state.unselectAll);
const saveComponent = useFlowsManagerStore((state) => state.saveComponent); const saveComponent = useFlowsManagerStore((state) => state.saveComponent);
const flows = useFlowsManagerStore((state) => state.flows); const flows = useFlowsManagerStore((state) => state.flows);
const version = useDarkStore((state) => state.version); const version = useDarkStore((state) => state.version);
@ -210,7 +209,8 @@ export default function NodeToolbarComponent({
useEffect(() => { useEffect(() => {
function onKeyDown(event: KeyboardEvent) { function onKeyDown(event: KeyboardEvent) {
if ( if (
(selected && isGroup) && selected &&
isGroup &&
(event.ctrlKey || event.metaKey) && (event.ctrlKey || event.metaKey) &&
event.key === "u" event.key === "u"
) { ) {
@ -234,10 +234,11 @@ export default function NodeToolbarComponent({
event.key === "S" event.key === "S"
) { ) {
event.preventDefault(); event.preventDefault();
setShowconfirmShare(state => !state); setShowconfirmShare((state) => !state);
} }
if ( if (
(selected && isMinimal) && selected &&
isMinimal &&
(event.ctrlKey || event.metaKey) && (event.ctrlKey || event.metaKey) &&
event.key === "q" event.key === "q"
) { ) {
@ -252,15 +253,16 @@ export default function NodeToolbarComponent({
event.key === "C" event.key === "C"
) { ) {
event.preventDefault(); event.preventDefault();
setOpenModal(state => !state); setOpenModal((state) => !state);
} }
if ( if (
(selected && !isGroup) && selected &&
!isGroup &&
(event.ctrlKey || event.metaKey) && (event.ctrlKey || event.metaKey) &&
event.key === "e" event.key === "e"
) { ) {
event.preventDefault(); event.preventDefault();
setShowModalAdvanced(state => !state); setShowModalAdvanced((state) => !state);
} }
if ( if (
selected && selected &&
@ -269,7 +271,7 @@ export default function NodeToolbarComponent({
isSaved isSaved
) { ) {
event.preventDefault(); event.preventDefault();
return setShowOverrideModal(state => !state); return setShowOverrideModal((state) => !state);
} }
if ( if (
selected && selected &&
@ -279,10 +281,11 @@ export default function NodeToolbarComponent({
) { ) {
event.preventDefault(); event.preventDefault();
saveComponent(cloneDeep(data), false); saveComponent(cloneDeep(data), false);
unselectAll() unselectAll();
} }
if ( if (
(selected && data.node?.documentation) && selected &&
data.node?.documentation &&
(event.ctrlKey || event.metaKey) && (event.ctrlKey || event.metaKey) &&
event.shiftKey && event.shiftKey &&
event.key === "D" event.key === "D"
@ -296,8 +299,8 @@ export default function NodeToolbarComponent({
return () => { return () => {
document.removeEventListener("keydown", onKeyDown); document.removeEventListener("keydown", onKeyDown);
} };
}, []) }, []);
return ( return (
<> <>
@ -313,7 +316,7 @@ export default function NodeToolbarComponent({
data-testid="code-button-modal" data-testid="code-button-modal"
> >
<div className="hidden"> <div className="hidden">
<CodeAreaComponent <CodeAreaComponent
openModal={openModal} openModal={openModal}
setOpenModal={setOpenModal} setOpenModal={setOpenModal}
readonly={ readonly={
@ -329,7 +332,7 @@ export default function NodeToolbarComponent({
onChange={handleOnNewValue} onChange={handleOnNewValue}
id={"code-input-node-toolbar-" + name} id={"code-input-node-toolbar-" + name}
selected={selected} selected={selected}
/> />
</div> </div>
<IconComponent name="TerminalSquare" className="h-4 w-4" /> <IconComponent name="TerminalSquare" className="h-4 w-4" />
</button> </button>
@ -430,50 +433,50 @@ export default function NodeToolbarComponent({
{nodeLength > 0 && ( {nodeLength > 0 && (
<SelectItem value={nodeLength === 0 ? "disabled" : "advanced"}> <SelectItem value={nodeLength === 0 ? "disabled" : "advanced"}>
<div className="flex"> <div className="flex">
<IconComponent <IconComponent
name="Settings2" name="Settings2"
className="relative top-0.5 mr-2 h-4 w-4 " className="relative top-0.5 mr-2 h-4 w-4 "
/>{" "} />{" "}
<span className="">Edit</span>{" "} <span className="">Edit</span>{" "}
<IconComponent <IconComponent
name="Command" name="Command"
className="absolute right-[1.15rem] top-[0.65em] h-3.5 w-3.5 stroke-2" className="absolute right-[1.15rem] top-[0.65em] h-3.5 w-3.5 stroke-2"
></IconComponent> ></IconComponent>
<span className="absolute right-2 top-[0.46em]">E</span> <span className="absolute right-2 top-[0.46em]">E</span>
</div> </div>
</SelectItem> </SelectItem>
)} )}
{isSaved ? ( {isSaved ? (
<SelectItem value={"override"}> <SelectItem value={"override"}>
<div className="flex"> <div className="flex">
<IconComponent <IconComponent
name="SaveAll" name="SaveAll"
className="relative top-0.5 mr-2 h-4 w-4 " className="relative top-0.5 mr-2 h-4 w-4 "
/>{" "} />{" "}
<span className="">Save</span>{" "} <span className="">Save</span>{" "}
<IconComponent <IconComponent
name="Command" name="Command"
className="absolute right-[1.15rem] top-[0.65em] h-3.5 w-3.5 stroke-2" className="absolute right-[1.15rem] top-[0.65em] h-3.5 w-3.5 stroke-2"
></IconComponent> ></IconComponent>
<span className="absolute right-2 top-[0.45em]">S</span> <span className="absolute right-2 top-[0.45em]">S</span>
</div> </div>
</SelectItem> </SelectItem>
) : ( ) : (
hasCode && ( hasCode && (
<SelectItem value={"SaveAll"}> <SelectItem value={"SaveAll"}>
<div className="flex"> <div className="flex">
<IconComponent <IconComponent
name="Settings2" name="Settings2"
className="relative top-0.5 mr-2 h-4 w-4 " className="relative top-0.5 mr-2 h-4 w-4 "
/>{" "} />{" "}
<span className="">Save</span>{" "} <span className="">Save</span>{" "}
<IconComponent <IconComponent
name="Command" name="Command"
className="absolute right-[1.15rem] top-[0.65em] h-3.5 w-3.5 stroke-2" className="absolute right-[1.15rem] top-[0.65em] h-3.5 w-3.5 stroke-2"
></IconComponent> ></IconComponent>
<span className="absolute right-2 top-[0.46em]">S</span> <span className="absolute right-2 top-[0.46em]">S</span>
</div> </div>
</SelectItem> </SelectItem>
) )
)} )}
@ -503,14 +506,14 @@ export default function NodeToolbarComponent({
/>{" "} />{" "}
Share{" "} Share{" "}
<IconComponent <IconComponent
name="Command" name="Command"
className="absolute right-[2rem] top-[0.65em] h-3.5 w-3.5 stroke-2" className="absolute right-[2rem] top-[0.65em] h-3.5 w-3.5 stroke-2"
></IconComponent> ></IconComponent>
<IconComponent <IconComponent
name="ArrowBigUp" name="ArrowBigUp"
className="absolute right-[1.09rem] top-[0.65em] h-3.5 w-3.5 stroke-2" className="absolute right-[1.09rem] top-[0.65em] h-3.5 w-3.5 stroke-2"
></IconComponent> ></IconComponent>
<span className="absolute right-2 top-[0.45em]">S</span> <span className="absolute right-2 top-[0.45em]">S</span>
</div>{" "} </div>{" "}
</SelectItem> </SelectItem>
)} )}
@ -549,33 +552,33 @@ export default function NodeToolbarComponent({
{isMinimal && ( {isMinimal && (
<SelectItem value={"show"}> <SelectItem value={"show"}>
<div className="flex"> <div className="flex">
<IconComponent <IconComponent
name={showNode ? "Minimize2" : "Maximize2"} name={showNode ? "Minimize2" : "Maximize2"}
className="relative top-0.5 mr-2 h-4 w-4" className="relative top-0.5 mr-2 h-4 w-4"
/> />
{showNode ? "Minimize" : "Expand"} {showNode ? "Minimize" : "Expand"}
<IconComponent <IconComponent
name="Command" name="Command"
className="absolute right-[1.25rem] top-[0.65em] h-3.5 w-3.5 stroke-2" className="absolute right-[1.25rem] top-[0.65em] h-3.5 w-3.5 stroke-2"
></IconComponent> ></IconComponent>
<span className="absolute right-2 top-[0.46em]">Q</span> <span className="absolute right-2 top-[0.46em]">Q</span>
</div> </div>
</SelectItem> </SelectItem>
)} )}
{isGroup && ( {isGroup && (
<SelectItem value="ungroup"> <SelectItem value="ungroup">
<div className="flex"> <div className="flex">
<IconComponent <IconComponent
name="Ungroup" name="Ungroup"
className="relative top-0.5 mr-2 h-4 w-4 " className="relative top-0.5 mr-2 h-4 w-4 "
/>{" "} />{" "}
<span className="">Ungroup</span>{" "} <span className="">Ungroup</span>{" "}
<IconComponent <IconComponent
name="Command" name="Command"
className="absolute right-[1.15rem] top-[0.65em] h-3.5 w-3.5 stroke-2" className="absolute right-[1.15rem] top-[0.65em] h-3.5 w-3.5 stroke-2"
></IconComponent> ></IconComponent>
<span className="absolute right-2 top-[0.43em]">U</span> <span className="absolute right-2 top-[0.43em]">U</span>
</div> </div>
</SelectItem> </SelectItem>
)} )}

View file

@ -219,8 +219,16 @@ const useFlowStore = create<FlowStoreType>((set, get) => ({
); );
}, },
paste: (selection, position) => { paste: (selection, position) => {
if(selection.nodes.some((node) => node.data.type === "ChatInput") && checkChatInput(get().nodes)){ if (
useAlertStore.getState().setErrorData({title: "Error pasting components", list: ["You can only have one ChatInput component in the flow"]}); selection.nodes.some((node) => node.data.type === "ChatInput") &&
checkChatInput(get().nodes)
) {
useAlertStore
.getState()
.setErrorData({
title: "Error pasting components",
list: ["You can only have one ChatInput component in the flow"],
});
return; return;
} }
let minimumX = Infinity; let minimumX = Infinity;

View file

@ -2,14 +2,14 @@ import { create } from "zustand";
import { shortcutsStoreType } from "../types/store"; import { shortcutsStoreType } from "../types/store";
export const useShortcutsStore = create<shortcutsStoreType>((set, get) => ({ export const useShortcutsStore = create<shortcutsStoreType>((set, get) => ({
openCodeModalWShortcut: false, openCodeModalWShortcut: false,
handleModalWShortcut: (modal) => { handleModalWShortcut: (modal) => {
switch (modal) { switch (modal) {
case "code": case "code":
set({ set({
openCodeModalWShortcut: !get().openCodeModalWShortcut, openCodeModalWShortcut: !get().openCodeModalWShortcut,
}); });
break; break;
} }
}, },
})); }));

View file

@ -22,4 +22,4 @@ export type StoreComponentResponse = {
export type shortcutsStoreType = { export type shortcutsStoreType = {
openCodeModalWShortcut: boolean; openCodeModalWShortcut: boolean;
handleModalWShortcut: (str: string) => void; handleModalWShortcut: (str: string) => void;
}; };

View file

@ -85,14 +85,12 @@ export async function updateVerticesOrder(
} }
const verticesIds = verticesLayers.flat(); const verticesIds = verticesLayers.flat();
useFlowStore useFlowStore.getState().updateVerticesBuild({
.getState() verticesLayers,
.updateVerticesBuild({ verticesIds,
verticesLayers, verticesOrder,
verticesIds, runId,
verticesOrder, });
runId,
});
resolve({ verticesLayers, verticesIds, verticesOrder, runId }); resolve({ verticesLayers, verticesIds, verticesOrder, runId });
}); });
} }

View file

@ -29,10 +29,7 @@ def poll_task_status(client, headers, href, max_attempts=20, sleep_time=1):
href, href,
headers=headers, headers=headers,
) )
if ( if task_status_response.status_code == 200 and task_status_response.json()["status"] == "SUCCESS":
task_status_response.status_code == 200
and task_status_response.json()["status"] == "SUCCESS"
):
return task_status_response.json() return task_status_response.json()
time.sleep(sleep_time) time.sleep(sleep_time)
return None # Return None if task did not complete in time return None # Return None if task did not complete in time
@ -126,11 +123,7 @@ def created_api_key(active_user):
) )
db_manager = get_db_service() db_manager = get_db_service()
with session_getter(db_manager) as session: with session_getter(db_manager) as session:
if ( if existing_api_key := session.query(ApiKey).filter(ApiKey.api_key == api_key.api_key).first():
existing_api_key := session.query(ApiKey)
.filter(ApiKey.api_key == api_key.api_key)
.first()
):
return existing_api_key return existing_api_key
session.add(api_key) session.add(api_key)
session.commit() session.commit()
@ -296,11 +289,7 @@ def test_get_all(client: TestClient, logged_in_headers):
dir_reader = DirectoryReader(settings.COMPONENTS_PATH[0]) dir_reader = DirectoryReader(settings.COMPONENTS_PATH[0])
files = dir_reader.get_files() files = dir_reader.get_files()
# json_response is a dict of dicts # json_response is a dict of dicts
all_names = [ all_names = [component_name for _, components in response.json().items() for component_name in components]
component_name
for _, components in response.json().items()
for component_name in components
]
json_response = response.json() json_response = response.json()
# We need to test the custom nodes # We need to test the custom nodes
assert len(all_names) > len(files) assert len(all_names) > len(files)
@ -425,19 +414,13 @@ def test_various_prompts(client, prompt, expected_input_variables):
def test_get_vertices_flow_not_found(client, logged_in_headers): def test_get_vertices_flow_not_found(client, logged_in_headers):
response = client.get( response = client.get("/api/v1/build/nonexistent_id/vertices", headers=logged_in_headers)
"/api/v1/build/nonexistent_id/vertices", headers=logged_in_headers assert response.status_code == 500 # Or whatever status code you've set for invalid ID
)
assert (
response.status_code == 500
) # Or whatever status code you've set for invalid ID
def test_get_vertices(client, added_flow_with_prompt_and_history, logged_in_headers): def test_get_vertices(client, added_flow_with_prompt_and_history, logged_in_headers):
flow_id = added_flow_with_prompt_and_history["id"] flow_id = added_flow_with_prompt_and_history["id"]
response = client.get( response = client.get(f"/api/v1/build/{flow_id}/vertices", headers=logged_in_headers)
f"/api/v1/build/{flow_id}/vertices", headers=logged_in_headers
)
assert response.status_code == 200 assert response.status_code == 200
assert "ids" in response.json() assert "ids" in response.json()
# The response should contain the list in this order # The response should contain the list in this order
@ -453,19 +436,13 @@ def test_get_vertices(client, added_flow_with_prompt_and_history, logged_in_head
def test_build_vertex_invalid_flow_id(client, logged_in_headers): def test_build_vertex_invalid_flow_id(client, logged_in_headers):
response = client.post( response = client.post("/api/v1/build/nonexistent_id/vertices/vertex_id", headers=logged_in_headers)
"/api/v1/build/nonexistent_id/vertices/vertex_id", headers=logged_in_headers
)
assert response.status_code == 500 assert response.status_code == 500
def test_build_vertex_invalid_vertex_id( def test_build_vertex_invalid_vertex_id(client, added_flow_with_prompt_and_history, logged_in_headers):
client, added_flow_with_prompt_and_history, logged_in_headers
):
flow_id = added_flow_with_prompt_and_history["id"] flow_id = added_flow_with_prompt_and_history["id"]
response = client.post( response = client.post(f"/api/v1/build/{flow_id}/vertices/invalid_vertex_id", headers=logged_in_headers)
f"/api/v1/build/{flow_id}/vertices/invalid_vertex_id", headers=logged_in_headers
)
assert response.status_code == 500 assert response.status_code == 500