Refactor TemplateField class in base.py
This commit is contained in:
parent
aa95686470
commit
3d645ee41d
2 changed files with 21 additions and 81 deletions
|
|
@ -72,17 +72,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
|
||||||
|
|
@ -110,9 +104,7 @@ class Vertex:
|
||||||
):
|
):
|
||||||
if edge.target_id not in edge_results:
|
if edge.target_id not in edge_results:
|
||||||
edge_results[edge.target_id] = {}
|
edge_results[edge.target_id] = {}
|
||||||
edge_results[edge.target_id][edge.target_param] = await edge.get_result(
|
edge_results[edge.target_id][edge.target_param] = await edge.get_result(source=self, target=target)
|
||||||
source=self, target=target
|
|
||||||
)
|
|
||||||
return edge_results
|
return edge_results
|
||||||
|
|
||||||
def set_result(self, result: "ResultData") -> None:
|
def set_result(self, result: "ResultData") -> None:
|
||||||
|
|
@ -122,9 +114,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
|
||||||
|
|
@ -134,11 +124,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
|
||||||
|
|
@ -201,29 +187,17 @@ class Vertex:
|
||||||
self.output = self.data["node"]["base_classes"]
|
self.output = self.data["node"]["base_classes"]
|
||||||
self.display_name = self.data["node"]["display_name"]
|
self.display_name = self.data["node"]["display_name"]
|
||||||
self.pinned = self.data["node"].get("pinned", False)
|
self.pinned = self.data["node"].get("pinned", False)
|
||||||
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.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"]
|
||||||
|
|
@ -266,11 +240,7 @@ class Vertex:
|
||||||
if self.graph is None:
|
if self.graph is None:
|
||||||
raise ValueError("Graph not found")
|
raise ValueError("Graph not found")
|
||||||
|
|
||||||
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:
|
||||||
|
|
@ -321,11 +291,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:
|
||||||
|
|
@ -388,9 +354,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):
|
||||||
|
|
@ -418,9 +382,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:
|
||||||
|
|
@ -454,9 +416,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.
|
||||||
"""
|
"""
|
||||||
|
|
@ -508,9 +468,7 @@ class Vertex:
|
||||||
self._update_built_object_and_artifacts(result)
|
self._update_built_object_and_artifacts(result)
|
||||||
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):
|
||||||
"""
|
"""
|
||||||
|
|
@ -580,24 +538,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:
|
||||||
|
|
@ -610,11 +560,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):
|
||||||
|
|
|
||||||
|
|
@ -68,9 +68,7 @@ class TemplateField(BaseModel):
|
||||||
refresh: Optional[bool] = None
|
refresh: Optional[bool] = None
|
||||||
"""Specifies if the field should be refreshed. Defaults to False."""
|
"""Specifies if the field should be refreshed. Defaults to False."""
|
||||||
|
|
||||||
range_spec: Optional[RangeSpec] = Field(
|
range_spec: Optional[RangeSpec] = Field(default=None, serialization_alias="rangeSpec")
|
||||||
default=None, serialization_alias="rangeSpec"
|
|
||||||
)
|
|
||||||
"""Range specification for the field. Defaults to None."""
|
"""Range specification for the field. Defaults to None."""
|
||||||
|
|
||||||
title_case: bool = False
|
title_case: bool = False
|
||||||
|
|
@ -117,10 +115,6 @@ class TemplateField(BaseModel):
|
||||||
if not isinstance(value, list):
|
if not isinstance(value, list):
|
||||||
raise ValueError("file_types must be a list")
|
raise ValueError("file_types must be a list")
|
||||||
return [
|
return [
|
||||||
(
|
(f".{file_type}" if isinstance(file_type, str) and not file_type.startswith(".") else file_type)
|
||||||
f".{file_type}"
|
|
||||||
if isinstance(file_type, str) and not file_type.startswith(".")
|
|
||||||
else file_type
|
|
||||||
)
|
|
||||||
for file_type in value
|
for file_type in value
|
||||||
]
|
]
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue