Add UnbuiltResult import and update vertex result type

This commit is contained in:
Gabriel Luiz Freitas Almeida 2024-02-29 23:08:21 -03:00
commit 1366967232
3 changed files with 14 additions and 12 deletions

View file

@ -18,6 +18,7 @@ from langflow.api.v1.schemas import (
VertexBuildResponse, VertexBuildResponse,
VerticesOrderResponse, VerticesOrderResponse,
) )
from langflow.graph.utils import UnbuiltResult
from langflow.services.auth.utils import get_current_active_user from langflow.services.auth.utils import get_current_active_user
from langflow.services.chat.service import ChatService from langflow.services.chat.service import ChatService
from langflow.services.deps import get_chat_service, get_session, get_session_service from langflow.services.deps import get_chat_service, get_session, get_session_service
@ -116,7 +117,7 @@ async def build_vertex(
inputs_dict = inputs.model_dump() if inputs else {} inputs_dict = inputs.model_dump() if inputs else {}
await vertex.build(user_id=current_user.id, inputs=inputs_dict) await vertex.build(user_id=current_user.id, inputs=inputs_dict)
if vertex.result is not None: if not isinstance(vertex.result, UnbuiltResult):
params = vertex._built_object_repr() params = vertex._built_object_repr()
valid = True valid = True
result_dict = vertex.result result_dict = vertex.result

View file

@ -10,7 +10,7 @@ from langflow.graph.graph.constants import lazy_load_vertex_dict
from langflow.graph.graph.state_manager import GraphStateManager from langflow.graph.graph.state_manager import GraphStateManager
from langflow.graph.graph.utils import process_flow from langflow.graph.graph.utils import process_flow
from langflow.graph.schema import INPUT_FIELD_NAME, InterfaceComponentTypes from langflow.graph.schema import INPUT_FIELD_NAME, InterfaceComponentTypes
from langflow.graph.vertex.base import Vertex from langflow.graph.vertex.base import Vertex, VertexStates
from langflow.graph.vertex.types import ( from langflow.graph.vertex.types import (
ChatVertex, ChatVertex,
FileToolVertex, FileToolVertex,
@ -148,17 +148,17 @@ class Graph:
def reset_inactive_vertices(self): def reset_inactive_vertices(self):
self.inactive_vertices = set() self.inactive_vertices = set()
def mark_all_vertices(self, state: str): def mark_all_vertices(self, state: "VertexStates"):
"""Marks all vertices in the graph.""" """Marks all vertices in the graph."""
for vertex in self.vertices: for vertex in self.vertices:
vertex.set_state(state) vertex.set_state(state)
def mark_vertex(self, vertex_id: str, state: str): def mark_vertex(self, vertex_id: str, state: "VertexStates"):
"""Marks a vertex in the graph.""" """Marks a vertex in the graph."""
vertex = self.get_vertex(vertex_id) vertex = self.get_vertex(vertex_id)
vertex.set_state(state) vertex.set_state(state)
def mark_branch(self, vertex_id: str, state: str): def mark_branch(self, vertex_id: str, state: "VertexStates"):
"""Marks a branch of the graph.""" """Marks a branch of the graph."""
self.mark_vertex(vertex_id, state) self.mark_vertex(vertex_id, state)
for child_id in self.parent_child_map[vertex_id]: for child_id in self.parent_child_map[vertex_id]:
@ -552,7 +552,7 @@ class Graph:
node_name = node_id.split("-")[0] node_name = node_id.split("-")[0]
if node_name in ["ChatOutput", "ChatInput"]: if node_name in ["ChatOutput", "ChatInput"]:
return ChatVertex return ChatVertex
elif node_name in ["ShouldRunNext"]: elif node_name in ["ShouldRunNext", "Branch"]:
return RoutingVertex return RoutingVertex
elif node_base_type in lazy_load_vertex_dict.VERTEX_TYPE_MAP: elif node_base_type in lazy_load_vertex_dict.VERTEX_TYPE_MAP:
return lazy_load_vertex_dict.VERTEX_TYPE_MAP[node_base_type] return lazy_load_vertex_dict.VERTEX_TYPE_MAP[node_base_type]

View file

@ -2,7 +2,7 @@ import ast
import inspect import inspect
import types import types
from enum import Enum from enum import Enum
from typing import TYPE_CHECKING, Any, Callable, Coroutine, Dict, List, Optional from typing import TYPE_CHECKING, Any, Callable, Coroutine, Dict, List, Optional, Union
from loguru import logger from loguru import logger
@ -75,7 +75,7 @@ class Vertex:
self.parent_is_top_level = False self.parent_is_top_level = False
self.layer = None self.layer = None
self.should_run = True self.should_run = True
self.result: Optional[ResultData] = None self.result: Union[ResultData, UnbuiltResult] = UnbuiltResult()
try: try:
self.is_interface_component = self.vertex_type in InterfaceComponentTypes self.is_interface_component = self.vertex_type in InterfaceComponentTypes
except ValueError: except ValueError:
@ -95,11 +95,12 @@ class Vertex:
else: else:
self.graph_state[key] = new_state self.graph_state[key] = new_state
def set_state(self, state: str): def set_state(self, state: "VertexStates"):
self.state = VertexStates[state] self.state = state
if ( if (
self.state == VertexStates.INACTIVE self.state
and self.graph.in_degree_map[self.id] < 2 == 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