Skip to content

Commit d3bd71a

Browse files
committed
fixing runtime and BaseNode
1 parent 80e9204 commit d3bd71a

2 files changed

Lines changed: 13 additions & 30 deletions

File tree

python-sdk/exospherehost/node/BaseNode.py

Lines changed: 3 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,5 @@
11
from abc import ABC, abstractmethod
2-
from typing import Any, Optional, List
2+
from typing import Optional, List
33
from pydantic import BaseModel
44

55

@@ -17,16 +17,15 @@ class BaseNode(ABC):
1717
state (dict[str, Any]): A dictionary for storing node state between executions.
1818
"""
1919

20-
def __init__(self, unique_name: Optional[str] = None):
20+
def __init__(self):
2121
"""
2222
Initialize a BaseNode instance.
2323
2424
Args:
2525
unique_name (Optional[str], optional): A unique identifier for this node.
2626
If None, the class name will be used as the unique name. Defaults to None.
2727
"""
28-
self.unique_name: Optional[str] = unique_name
29-
self.state: dict[str, Any] = {}
28+
self.inputs: Optional[BaseNode.Inputs] = None
3029

3130
class Inputs(BaseModel):
3231
"""

python-sdk/exospherehost/runtime.py

Lines changed: 10 additions & 26 deletions
Original file line numberDiff line numberDiff line change
@@ -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

Comments
 (0)