Refactor RoutingVertex to handle missing condition and result values
This commit is contained in:
parent
dcf71b7bdf
commit
75ee16acb0
1 changed files with 7 additions and 24 deletions
|
|
@ -373,35 +373,18 @@ class RoutingVertex(StatelessVertex):
|
||||||
return self.artifacts["repr"] or super()._built_object_repr()
|
return self.artifacts["repr"] or super()._built_object_repr()
|
||||||
return super()._built_object_repr()
|
return super()._built_object_repr()
|
||||||
|
|
||||||
def _build(self, *args, **kwargs):
|
|
||||||
super()._build(*args, **kwargs)
|
|
||||||
|
|
||||||
# After building, the _built_object should be a dict with
|
|
||||||
# {"result": Any, "condition": bool}
|
|
||||||
# if true, we need to set should_run attr in the target of true edge
|
|
||||||
# to true and should_run attr in the target of false edge to false
|
|
||||||
# TODO: Add support for multiple conditions
|
|
||||||
|
|
||||||
def _run(self, *args, **kwargs):
|
def _run(self, *args, **kwargs):
|
||||||
if self._built_object:
|
if self._built_object:
|
||||||
condition = self._built_object.get("condition")
|
condition = self._built_object.get("condition")
|
||||||
result = self._built_object.get("result")
|
result = self._built_object.get("result")
|
||||||
if condition is not None:
|
if condition is None:
|
||||||
for edge in self.edges:
|
raise ValueError("Condition is required for the routing vertex.")
|
||||||
if edge.source_id == self.id:
|
if result is None:
|
||||||
target_vertex = self.graph.get_vertex(edge.target_id)
|
raise ValueError("Result is required for the routing vertex.")
|
||||||
# source_handle.channel and condition should be the same
|
if condition is True:
|
||||||
channel_bool = edge.source_handle.channel == "true"
|
self._built_result = result
|
||||||
if condition == channel_bool:
|
|
||||||
target_vertex.should_run = True
|
|
||||||
else:
|
|
||||||
target_vertex.should_run = False
|
|
||||||
else:
|
else:
|
||||||
raise ValueError(f"RoutingVertex {self.id} must have a condition in the _built_object")
|
self.graph.mark_branch(self.id, "INACTIVE")
|
||||||
|
|
||||||
self._built_result = result
|
|
||||||
else:
|
|
||||||
raise ValueError(f"RoutingVertex {self.id} must have a _built_object with a condition and a result")
|
|
||||||
|
|
||||||
|
|
||||||
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