Change honor method to be asynchronous in ContractEdge
This commit is contained in:
parent
b9b3ef88fe
commit
20d9e51208
3 changed files with 38 additions and 23 deletions
|
|
@ -105,7 +105,7 @@ class ContractEdge(Edge):
|
||||||
self.is_fulfilled = False # Whether the contract has been fulfilled.
|
self.is_fulfilled = False # Whether the contract has been fulfilled.
|
||||||
self.result: Any = None
|
self.result: Any = None
|
||||||
|
|
||||||
def honor(self, source: "Vertex", target: "Vertex") -> None:
|
async def honor(self, source: "Vertex", target: "Vertex") -> None:
|
||||||
"""
|
"""
|
||||||
Fulfills the contract by setting the result of the source vertex to the target vertex's parameter.
|
Fulfills the contract by setting the result of the source vertex to the target vertex's parameter.
|
||||||
If the edge is runnable, the source vertex is run with the message text and the target vertex's
|
If the edge is runnable, the source vertex is run with the message text and the target vertex's
|
||||||
|
|
@ -117,7 +117,7 @@ class ContractEdge(Edge):
|
||||||
return
|
return
|
||||||
|
|
||||||
if not source._built:
|
if not source._built:
|
||||||
source.build()
|
await source.build()
|
||||||
|
|
||||||
if self.matched_type == "Text":
|
if self.matched_type == "Text":
|
||||||
self.result = source._built_result
|
self.result = source._built_result
|
||||||
|
|
@ -144,7 +144,7 @@ class ContractEdge(Edge):
|
||||||
async def get_result(self, source: "Vertex", target: "Vertex"):
|
async def get_result(self, source: "Vertex", target: "Vertex"):
|
||||||
# Fulfill the contract if it has not been fulfilled.
|
# Fulfill the contract if it has not been fulfilled.
|
||||||
if not self.is_fulfilled:
|
if not self.is_fulfilled:
|
||||||
self.honor(source, target)
|
await self.honor(source, target)
|
||||||
|
|
||||||
log_transaction(self, source, target, "success")
|
log_transaction(self, source, target, "success")
|
||||||
# If the target vertex is a power component we log messages
|
# If the target vertex is a power component we log messages
|
||||||
|
|
|
||||||
|
|
@ -4,7 +4,7 @@ from typing import Dict, Generator, List, Type, Union
|
||||||
from langchain.chains.base import Chain
|
from langchain.chains.base import Chain
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
|
|
||||||
from langflow.graph.edge.base import ContractEdge, Edge
|
from langflow.graph.edge.base import ContractEdge
|
||||||
from langflow.graph.graph.constants import lazy_load_vertex_dict
|
from langflow.graph.graph.constants import lazy_load_vertex_dict
|
||||||
from langflow.graph.graph.utils import process_flow
|
from langflow.graph.graph.utils import process_flow
|
||||||
from langflow.graph.vertex.base import Vertex
|
from langflow.graph.vertex.base import Vertex
|
||||||
|
|
@ -33,6 +33,7 @@ class Graph:
|
||||||
|
|
||||||
self._vertices = self._graph_data["nodes"]
|
self._vertices = self._graph_data["nodes"]
|
||||||
self._edges = self._graph_data["edges"]
|
self._edges = self._graph_data["edges"]
|
||||||
|
|
||||||
self._build_graph()
|
self._build_graph()
|
||||||
|
|
||||||
def __getstate__(self):
|
def __getstate__(self):
|
||||||
|
|
@ -111,7 +112,7 @@ class Graph:
|
||||||
"""Returns a vertex by id."""
|
"""Returns a vertex by id."""
|
||||||
return self.vertex_map.get(vertex_id)
|
return self.vertex_map.get(vertex_id)
|
||||||
|
|
||||||
def get_vertex_edges(self, vertex_id: str) -> List[Union[Edge, ContractEdge]]:
|
def get_vertex_edges(self, vertex_id: str) -> List[ContractEdge]:
|
||||||
"""Returns a list of edges for a given vertex."""
|
"""Returns a list of edges for a given vertex."""
|
||||||
return [edge for edge in self.edges if edge.source_id == vertex_id or edge.target_id == vertex_id]
|
return [edge for edge in self.edges if edge.source_id == vertex_id or edge.target_id == vertex_id]
|
||||||
|
|
||||||
|
|
@ -210,18 +211,19 @@ class Graph:
|
||||||
edges.append(ContractEdge(source, target, edge))
|
edges.append(ContractEdge(source, target, edge))
|
||||||
return edges
|
return edges
|
||||||
|
|
||||||
def _get_vertex_class(self, vertex_type: str, vertex_base_type: str) -> Type[Vertex]:
|
def _get_vertex_class(self, node_type: str, node_lc_type: str, node_id: str) -> Type[Vertex]:
|
||||||
"""Returns the vertex class based on the vertex type."""
|
"""Returns the node class based on the node type."""
|
||||||
if vertex_type in FILE_TOOLS:
|
node_name = node_id.split("-")[0]
|
||||||
return FileToolVertex
|
if node_name in lazy_load_vertex_dict.VERTEX_TYPE_MAP:
|
||||||
if vertex_base_type == "CustomComponent":
|
return lazy_load_vertex_dict.VERTEX_TYPE_MAP[node_name]
|
||||||
return lazy_load_vertex_dict.get_custom_component_vertex_type()
|
|
||||||
|
|
||||||
if vertex_base_type in lazy_load_vertex_dict.VERTEX_TYPE_MAP:
|
if node_type in FILE_TOOLS:
|
||||||
return lazy_load_vertex_dict.VERTEX_TYPE_MAP[vertex_base_type]
|
return FileToolVertex
|
||||||
|
if node_type in lazy_load_vertex_dict.VERTEX_TYPE_MAP:
|
||||||
|
return lazy_load_vertex_dict.VERTEX_TYPE_MAP[node_type]
|
||||||
return (
|
return (
|
||||||
lazy_load_vertex_dict.VERTEX_TYPE_MAP[vertex_type]
|
lazy_load_vertex_dict.VERTEX_TYPE_MAP[node_lc_type]
|
||||||
if vertex_type in lazy_load_vertex_dict.VERTEX_TYPE_MAP
|
if node_lc_type in lazy_load_vertex_dict.VERTEX_TYPE_MAP
|
||||||
else Vertex
|
else Vertex
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
@ -233,7 +235,7 @@ 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(vertex_type, vertex_base_type)
|
VertexClass = self._get_vertex_class(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)
|
||||||
|
|
|
||||||
|
|
@ -1,6 +1,8 @@
|
||||||
import ast
|
import ast
|
||||||
from typing import Callable, Dict, List, Optional, Union
|
from typing import Callable, Dict, List, Optional, Union
|
||||||
|
|
||||||
|
from langchain_core.messages import AIMessage
|
||||||
|
|
||||||
from langflow.graph.utils import UnbuiltObject, flatten_list
|
from langflow.graph.utils import UnbuiltObject, flatten_list
|
||||||
from langflow.graph.vertex.base import StatefulVertex, StatelessVertex
|
from langflow.graph.vertex.base import StatefulVertex, StatelessVertex
|
||||||
from langflow.interface.utils import extract_input_variables_from_prompt
|
from langflow.interface.utils import extract_input_variables_from_prompt
|
||||||
|
|
@ -318,17 +320,28 @@ class ChatVertex(StatelessVertex):
|
||||||
if self.artifacts and "repr" in self.artifacts:
|
if self.artifacts and "repr" in self.artifacts:
|
||||||
return self.artifacts["repr"] or super()._built_object_repr()
|
return self.artifacts["repr"] or super()._built_object_repr()
|
||||||
|
|
||||||
def _run(self, *args, **kwargs):
|
async def _run(self, *args, **kwargs):
|
||||||
if self.is_power_component:
|
if self.is_power_component:
|
||||||
if self.vertex_type == "ChatOutput":
|
if self.vertex_type == "ChatOutput":
|
||||||
|
artifacts = None
|
||||||
sender = self.params.get("sender", None)
|
sender = self.params.get("sender", None)
|
||||||
sender_name = self.params.get("sender_name", None)
|
sender_name = self.params.get("sender_name", None)
|
||||||
self.artifacts = ChatOutputResponse(
|
message = ""
|
||||||
message=str(self._built_object),
|
if isinstance(self._built_object, AIMessage):
|
||||||
sender=sender,
|
artifacts = ChatOutputResponse.from_message(
|
||||||
sender_name=sender_name,
|
self._built_object,
|
||||||
).model_dump()
|
sender=sender,
|
||||||
|
sender_name=sender_name,
|
||||||
|
)
|
||||||
|
elif not isinstance(self._built_object, UnbuiltObject):
|
||||||
|
artifacts = ChatOutputResponse(
|
||||||
|
message=message,
|
||||||
|
sender=sender,
|
||||||
|
sender_name=sender_name,
|
||||||
|
)
|
||||||
|
if artifacts:
|
||||||
|
self.artifacts = artifacts.model_dump()
|
||||||
self._built_result = self._built_object
|
self._built_result = self._built_object
|
||||||
|
|
||||||
else:
|
else:
|
||||||
super()._run(*args, **kwargs)
|
await super()._run(*args, **kwargs)
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue