@@ -54,16 +54,18 @@ def __init__(self, namespace: str, name: str, nodes: List[type[BaseNode]], state
5454 self ._batch_size = batch_size
5555 self ._state_queue = Queue (maxsize = 2 * batch_size )
5656 self ._workers = workers
57- self ._nodes = []
58- self ._node_names = []
57+ self ._nodes = nodes
58+ self ._node_names = [node . __class__ . __name__ for node in nodes ]
5959 self ._state_manager_uri = state_manager_uri
6060 self ._state_manager_version = state_manage_version
6161 self ._poll_interval = poll_interval
62- self ._node_mapping = {}
62+ self ._node_mapping = {
63+ node .__class__ .__name__ : node for node in nodes
64+ }
6365
6466 self ._set_config_from_env ()
6567 self ._validate_runtime ()
66- self ._validate_nodes (nodes )
68+ self ._validate_nodes ()
6769
6870 def _set_config_from_env (self ):
6971 """
@@ -115,7 +117,7 @@ def _get_register_endpoint(self):
115117 """
116118 return f"{ self ._state_manager_uri } /{ str (self ._state_manager_version )} /namespace/{ self ._namespace } /nodes/"
117119
118- async def _register_nodes (self ):
120+ async def _register (self ):
119121 """
120122 Register node schemas and runtime metadata with the state manager.
121123
@@ -146,22 +148,6 @@ async def _register_nodes(self):
146148
147149 return res
148150
149- async def _register (self , nodes : List [type [BaseNode ]]):
150- """
151- Validate and register nodes with the runtime and state manager.
152-
153- Args:
154- nodes (List[type[BaseNode]]): List of BaseNode subclasses to register.
155-
156- Raises:
157- ValidationError: If any node is invalid.
158- """
159- self ._nodes = self ._validate_nodes (nodes )
160- self ._node_names = [node .__class__ .__name__ for node in nodes ]
161- self ._node_mapping = {node .__class__ .__name__ : node for node in self ._nodes }
162-
163- await self ._register_nodes ()
164-
165151 async def _enqueue_call (self ):
166152 """
167153 Request a batch of states to process from the state manager.
@@ -237,7 +223,7 @@ async def _notify_errored(self, state_id: str, error: str):
237223 if response .status != 200 :
238224 logger .error (f"Failed to notify errored state { state_id } : { res } " )
239225
240- def _validate_nodes (self , nodes : List [ type [ BaseNode ]] ):
226+ def _validate_nodes (self ):
241227 """
242228 Validate that all provided nodes are valid BaseNode subclasses.
243229
@@ -252,7 +238,7 @@ def _validate_nodes(self, nodes: List[type[BaseNode]]):
252238 """
253239 errors = []
254240
255- for node in nodes :
241+ for node in self . _nodes :
256242 if not issubclass (node , BaseNode ):
257243 errors .append (f"{ node .__class__ .__name__ } does not inherit from exospherehost.BaseNode" )
258244 if not hasattr (node , "Inputs" ):
@@ -265,16 +251,14 @@ def _validate_nodes(self, nodes: List[type[BaseNode]]):
265251 errors .append (f"{ node .__class__ .__name__ } does not have an Outputs class that inherits from pydantic.BaseModel" )
266252
267253 # Find nodes with the same __class__.__name__
268- class_names = [node .__class__ .__name__ for node in nodes ]
254+ class_names = [node .__class__ .__name__ for node in self . _nodes ]
269255 duplicate_class_names = [name for name in set (class_names ) if class_names .count (name ) > 1 ]
270256 if duplicate_class_names :
271257 errors .append (f"Duplicate node class names found: { duplicate_class_names } " )
272258
273259 if len (errors ) > 0 :
274260 raise ValidationError ("Following errors while validating nodes: " + "\n " .join (errors ))
275261
276- return nodes
277-
278262 async def _worker (self ):
279263 """
280264 Worker task that processes states from the queue.
0 commit comments