Add branches functionality to RoutingVertex
This commit is contained in:
parent
1a144a2fc1
commit
803d51ac25
1 changed files with 86 additions and 24 deletions
|
|
@ -1,5 +1,6 @@
|
||||||
import ast
|
import ast
|
||||||
import json
|
import json
|
||||||
|
from collections import defaultdict
|
||||||
from typing import AsyncIterator, Callable, Dict, Iterator, List, Optional, Union
|
from typing import AsyncIterator, Callable, Dict, Iterator, List, Optional, Union
|
||||||
|
|
||||||
import yaml
|
import yaml
|
||||||
|
|
@ -8,7 +9,7 @@ from loguru import logger
|
||||||
|
|
||||||
from langflow.graph.schema import INPUT_FIELD_NAME
|
from langflow.graph.schema import INPUT_FIELD_NAME
|
||||||
from langflow.graph.utils import UnbuiltObject, flatten_list, serialize_field
|
from langflow.graph.utils import UnbuiltObject, flatten_list, serialize_field
|
||||||
from langflow.graph.vertex.base import StatefulVertex, StatelessVertex
|
from langflow.graph.vertex.base import StatefulVertex, StatelessVertex, VertexStates
|
||||||
from langflow.interface.utils import extract_input_variables_from_prompt
|
from langflow.interface.utils import extract_input_variables_from_prompt
|
||||||
from langflow.schema import Record
|
from langflow.schema import Record
|
||||||
from langflow.services.monitor.utils import log_vertex_build
|
from langflow.services.monitor.utils import log_vertex_build
|
||||||
|
|
@ -123,9 +124,11 @@ class DocumentLoaderVertex(StatefulVertex):
|
||||||
# show how many documents are in the list?
|
# show how many documents are in the list?
|
||||||
|
|
||||||
if not isinstance(self._built_object, UnbuiltObject):
|
if not isinstance(self._built_object, UnbuiltObject):
|
||||||
avg_length = sum(len(doc.page_content) for doc in self._built_object if hasattr(doc, "page_content")) / len(
|
avg_length = sum(
|
||||||
self._built_object
|
len(doc.page_content)
|
||||||
)
|
for doc in self._built_object
|
||||||
|
if hasattr(doc, "page_content")
|
||||||
|
) / len(self._built_object)
|
||||||
return f"""{self.display_name}({len(self._built_object)} documents)
|
return f"""{self.display_name}({len(self._built_object)} documents)
|
||||||
\nAvg. Document Length (characters): {int(avg_length)}
|
\nAvg. Document Length (characters): {int(avg_length)}
|
||||||
Documents: {self._built_object[:3]}..."""
|
Documents: {self._built_object[:3]}..."""
|
||||||
|
|
@ -198,7 +201,9 @@ class TextSplitterVertex(StatefulVertex):
|
||||||
# show how many documents are in the list?
|
# show how many documents are in the list?
|
||||||
|
|
||||||
if not isinstance(self._built_object, UnbuiltObject):
|
if not isinstance(self._built_object, UnbuiltObject):
|
||||||
avg_length = sum(len(doc.page_content) for doc in self._built_object) / len(self._built_object)
|
avg_length = sum(len(doc.page_content) for doc in self._built_object) / len(
|
||||||
|
self._built_object
|
||||||
|
)
|
||||||
return f"""{self.vertex_type}({len(self._built_object)} documents)
|
return f"""{self.vertex_type}({len(self._built_object)} documents)
|
||||||
\nAvg. Document Length (characters): {int(avg_length)}
|
\nAvg. Document Length (characters): {int(avg_length)}
|
||||||
\nDocuments: {self._built_object[:3]}..."""
|
\nDocuments: {self._built_object[:3]}..."""
|
||||||
|
|
@ -245,18 +250,27 @@ class PromptVertex(StatelessVertex):
|
||||||
user_id = kwargs.get("user_id", None)
|
user_id = kwargs.get("user_id", None)
|
||||||
tools = kwargs.get("tools", [])
|
tools = kwargs.get("tools", [])
|
||||||
if not self._built or force:
|
if not self._built or force:
|
||||||
if "input_variables" not in self.params or self.params["input_variables"] is None:
|
if (
|
||||||
|
"input_variables" not in self.params
|
||||||
|
or self.params["input_variables"] is None
|
||||||
|
):
|
||||||
self.params["input_variables"] = []
|
self.params["input_variables"] = []
|
||||||
# Check if it is a ZeroShotPrompt and needs a tool
|
# Check if it is a ZeroShotPrompt and needs a tool
|
||||||
if "ShotPrompt" in self.vertex_type:
|
if "ShotPrompt" in self.vertex_type:
|
||||||
tools = [tool_node.build(user_id=user_id) for tool_node in tools] if tools is not None else []
|
tools = (
|
||||||
|
[tool_node.build(user_id=user_id) for tool_node in tools]
|
||||||
|
if tools is not None
|
||||||
|
else []
|
||||||
|
)
|
||||||
# flatten the list of tools if it is a list of lists
|
# flatten the list of tools if it is a list of lists
|
||||||
# first check if it is a list
|
# first check if it is a list
|
||||||
if tools and isinstance(tools, list) and isinstance(tools[0], list):
|
if tools and isinstance(tools, list) and isinstance(tools[0], list):
|
||||||
tools = flatten_list(tools)
|
tools = flatten_list(tools)
|
||||||
self.params["tools"] = tools
|
self.params["tools"] = tools
|
||||||
prompt_params = [
|
prompt_params = [
|
||||||
key for key, value in self.params.items() if isinstance(value, str) and key != "format_instructions"
|
key
|
||||||
|
for key, value in self.params.items()
|
||||||
|
if isinstance(value, str) and key != "format_instructions"
|
||||||
]
|
]
|
||||||
else:
|
else:
|
||||||
prompt_params = ["template"]
|
prompt_params = ["template"]
|
||||||
|
|
@ -266,14 +280,20 @@ class PromptVertex(StatelessVertex):
|
||||||
prompt_text = self.params[param]
|
prompt_text = self.params[param]
|
||||||
variables = extract_input_variables_from_prompt(prompt_text)
|
variables = extract_input_variables_from_prompt(prompt_text)
|
||||||
self.params["input_variables"].extend(variables)
|
self.params["input_variables"].extend(variables)
|
||||||
self.params["input_variables"] = list(set(self.params["input_variables"]))
|
self.params["input_variables"] = list(
|
||||||
|
set(self.params["input_variables"])
|
||||||
|
)
|
||||||
elif isinstance(self.params, dict):
|
elif isinstance(self.params, dict):
|
||||||
self.params.pop("input_variables", None)
|
self.params.pop("input_variables", None)
|
||||||
|
|
||||||
await self._build(user_id=user_id)
|
await self._build(user_id=user_id)
|
||||||
|
|
||||||
def _built_object_repr(self):
|
def _built_object_repr(self):
|
||||||
if not self.artifacts or self._built_object is None or not hasattr(self._built_object, "format"):
|
if (
|
||||||
|
not self.artifacts
|
||||||
|
or self._built_object is None
|
||||||
|
or not hasattr(self._built_object, "format")
|
||||||
|
):
|
||||||
return super()._built_object_repr()
|
return super()._built_object_repr()
|
||||||
elif isinstance(self._built_object, UnbuiltObject):
|
elif isinstance(self._built_object, UnbuiltObject):
|
||||||
return super()._built_object_repr()
|
return super()._built_object_repr()
|
||||||
|
|
@ -285,7 +305,9 @@ class PromptVertex(StatelessVertex):
|
||||||
# so the prompt format doesn't break
|
# so the prompt format doesn't break
|
||||||
artifacts.pop("handle_keys", None)
|
artifacts.pop("handle_keys", None)
|
||||||
try:
|
try:
|
||||||
if not hasattr(self._built_object, "template") and hasattr(self._built_object, "prompt"):
|
if not hasattr(self._built_object, "template") and hasattr(
|
||||||
|
self._built_object, "prompt"
|
||||||
|
):
|
||||||
template = self._built_object.prompt.template
|
template = self._built_object.prompt.template
|
||||||
else:
|
else:
|
||||||
template = self._built_object.template
|
template = self._built_object.template
|
||||||
|
|
@ -293,7 +315,11 @@ class PromptVertex(StatelessVertex):
|
||||||
if value:
|
if value:
|
||||||
replace_key = "{" + key + "}"
|
replace_key = "{" + key + "}"
|
||||||
template = template.replace(replace_key, value)
|
template = template.replace(replace_key, value)
|
||||||
return template if isinstance(template, str) else f"{self.vertex_type}({template})"
|
return (
|
||||||
|
template
|
||||||
|
if isinstance(template, str)
|
||||||
|
else f"{self.vertex_type}({template})"
|
||||||
|
)
|
||||||
except KeyError:
|
except KeyError:
|
||||||
return str(self._built_object)
|
return str(self._built_object)
|
||||||
|
|
||||||
|
|
@ -436,6 +462,24 @@ class RoutingVertex(StatelessVertex):
|
||||||
super().__init__(data, graph=graph, base_type="custom_components")
|
super().__init__(data, graph=graph, base_type="custom_components")
|
||||||
self.use_result = True
|
self.use_result = True
|
||||||
self.steps = [self._build, self._run]
|
self.steps = [self._build, self._run]
|
||||||
|
self._branches = defaultdict(set)
|
||||||
|
|
||||||
|
def build_branches(self):
|
||||||
|
if self._branches:
|
||||||
|
return
|
||||||
|
for edge in self.edges:
|
||||||
|
if edge.target_id == self.id:
|
||||||
|
continue
|
||||||
|
if edge.source_handle.conditionalPath is not None:
|
||||||
|
self._branches[edge.source_handle.conditionalPath].add(edge.target_id)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def true_branch(self):
|
||||||
|
return self._branches.get(True, set())
|
||||||
|
|
||||||
|
@property
|
||||||
|
def false_branch(self):
|
||||||
|
return self._branches.get(False, set())
|
||||||
|
|
||||||
def _built_object_repr(self):
|
def _built_object_repr(self):
|
||||||
if self.artifacts and "repr" in self.artifacts:
|
if self.artifacts and "repr" in self.artifacts:
|
||||||
|
|
@ -443,18 +487,36 @@ class RoutingVertex(StatelessVertex):
|
||||||
return super()._built_object_repr()
|
return super()._built_object_repr()
|
||||||
|
|
||||||
def _run(self, *args, **kwargs):
|
def _run(self, *args, **kwargs):
|
||||||
if self._built_object:
|
|
||||||
condition = self._built_object.get("condition")
|
self.build_branches()
|
||||||
result = self._built_object.get("result")
|
condition_path = self._built_object.get("path")
|
||||||
if condition is None:
|
result = self._built_object.get("result")
|
||||||
raise ValueError("Condition is required for the routing vertex.")
|
try:
|
||||||
if result is None:
|
# check if bool
|
||||||
raise ValueError("Result is required for the routing vertex.")
|
condition_path = bool(condition_path)
|
||||||
if condition is True:
|
except ValueError:
|
||||||
self._built_result = result
|
raise ValueError("'path' must be a boolean value.")
|
||||||
else:
|
|
||||||
self.graph.mark_branch(self.id, "INACTIVE")
|
# Validate necessary components are present
|
||||||
self._built_result = None
|
if not isinstance(condition_path, bool):
|
||||||
|
raise ValueError("Condition is required for the routing vertex.")
|
||||||
|
if result is None:
|
||||||
|
raise ValueError("Result is required for the routing vertex.")
|
||||||
|
if self._branches:
|
||||||
|
# Deactivate the branch not taken
|
||||||
|
self._deactivate_branch(condition_path)
|
||||||
|
self._built_result = result
|
||||||
|
elif condition_path is True:
|
||||||
|
self._built_result = result
|
||||||
|
else:
|
||||||
|
self.graph.mark_branch(self.id, VertexStates.INACTIVE)
|
||||||
|
|
||||||
|
def _deactivate_branch(self, condition_path: bool):
|
||||||
|
"""Deactivates the branch not taken based on the condition."""
|
||||||
|
# self.graph.mark_branch(target_id, "INACTIVE")
|
||||||
|
branch_to_deactivate = self.false_branch if condition_path else self.true_branch
|
||||||
|
for target_id in branch_to_deactivate:
|
||||||
|
self.graph.mark_branch(target_id, VertexStates.INACTIVE)
|
||||||
|
|
||||||
|
|
||||||
def dict_to_codeblock(d: dict) -> str:
|
def dict_to_codeblock(d: dict) -> str:
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue