33from app .models .db .registered_node import RegisteredNode
44from app .singletons .logs_manager import LogsManager
55from beanie .operators import In
6+ from json_schema_to_pydantic import create_model
7+ from collections import deque
68
79logger = LogsManager ().get_logger ()
810
@@ -16,26 +18,11 @@ async def verify_nodes_namespace(nodes: list[NodeTemplate], graph_namespace: str
1618 if node .namespace != graph_namespace and node .namespace != "exospherehost" :
1719 errors .append (f"Node { node .identifier } has invalid namespace '{ node .namespace } '. Must match graph namespace '{ graph_namespace } ' or use universal namespace 'exospherehost'" )
1820
19- async def verify_node_exists (nodes : list [NodeTemplate ], graph_namespace : str , errors : list [str ]):
20- graph_namespace_node_names = [
21- node .node_name for node in nodes if node .namespace == graph_namespace
22- ]
23- graph_namespace_database_nodes = await RegisteredNode .find (
24- In (RegisteredNode .name , graph_namespace_node_names ),
25- RegisteredNode .namespace == graph_namespace
26- ).to_list ()
27- exospherehost_node_names = [
28- node .node_name for node in nodes if node .namespace == "exospherehost"
29- ]
30- exospherehost_database_nodes = await RegisteredNode .find (
31- In (RegisteredNode .name , exospherehost_node_names ),
32- RegisteredNode .namespace == "exospherehost"
33- ).to_list ()
34-
35- template_nodes = set ([(node .node_name , node .namespace ) for node in nodes ])
36- database_nodes = set ([(node .name , node .namespace ) for node in graph_namespace_database_nodes + exospherehost_database_nodes ])
21+ async def verify_node_exists (nodes : list [NodeTemplate ], database_nodes : list [RegisteredNode ], errors : list [str ]):
22+ template_nodes_set = set ([(node .node_name , node .namespace ) for node in nodes ])
23+ database_nodes_set = set ([(node .name , node .namespace ) for node in database_nodes ])
3724
38- nodes_not_found = template_nodes - database_nodes
25+ nodes_not_found = template_nodes_set - database_nodes_set
3926
4027 for node in nodes_not_found :
4128 errors .append (f"Node { node [0 ]} in namespace { node [1 ]} does not exist." )
@@ -68,20 +55,179 @@ async def verify_node_identifiers(nodes: list[NodeTemplate], errors: list[str]):
6855 if next_node not in valid_identifiers :
6956 errors .append (f"Node { node .node_name } in namespace { node .namespace } has a next node { next_node } that does not exist in the graph" )
7057
58+ async def verify_secrets (graph_template : GraphTemplate , database_nodes : list [RegisteredNode ], errors : list [str ]):
59+ required_secrets_set = set ()
60+
61+ for node in database_nodes :
62+ if node .secrets is None :
63+ continue
64+ for secret in node .secrets :
65+ required_secrets_set .add (secret )
66+
67+ present_secrets_set = set ()
68+ for secret_name in graph_template .secrets .keys ():
69+ present_secrets_set .add (secret_name )
70+
71+ missing_secrets_set = required_secrets_set - present_secrets_set
72+
73+ for secret_name in missing_secrets_set :
74+ errors .append (f"Secret { secret_name } is required but not present in the graph template" )
75+
76+
77+ async def get_database_nodes (nodes : list [NodeTemplate ], graph_namespace : str ):
78+ graph_namespace_node_names = [
79+ node .node_name for node in nodes if node .namespace == graph_namespace
80+ ]
81+ graph_namespace_database_nodes = await RegisteredNode .find (
82+ In (RegisteredNode .name , graph_namespace_node_names ),
83+ RegisteredNode .namespace == graph_namespace
84+ ).to_list ()
85+ exospherehost_node_names = [
86+ node .node_name for node in nodes if node .namespace == "exospherehost"
87+ ]
88+ exospherehost_database_nodes = await RegisteredNode .find (
89+ In (RegisteredNode .name , exospherehost_node_names ),
90+ RegisteredNode .namespace == "exospherehost"
91+ ).to_list ()
92+ return graph_namespace_database_nodes + exospherehost_database_nodes
93+
94+
95+ async def verify_inputs (graph_nodes : list [NodeTemplate ], database_nodes : list [RegisteredNode ], dependencies_graph : dict [str , set [str ]], errors : list [str ]):
96+ look_up_table = {}
97+ for node in graph_nodes :
98+ look_up_table [node .identifier ] = {"graph_node" : node }
99+
100+ for database_node in database_nodes :
101+ if database_node .name == node .node_name and database_node .namespace == node .namespace :
102+ look_up_table [node .identifier ]["database_node" ] = database_node
103+ break
104+
105+ for node in graph_nodes :
106+ try :
107+ model_class = create_model (look_up_table [node .identifier ]["database_node" ].inputs_schema )
108+
109+ for field_name , field_info in model_class .model_fields .items ():
110+ if field_info .annotation is not str :
111+ errors .append (f"{ node .node_name } .Inputs field '{ field_name } ' must be of type str, got { field_info .annotation } " )
112+ continue
113+
114+ if field_name not in look_up_table [node .identifier ]["graph_node" ].inputs .keys ():
115+ errors .append (f"{ node .node_name } .Inputs field '{ field_name } ' not found in graph template" )
116+ continue
117+
118+ # get ${{ identifier.outputs.field_name }} objects from the string
119+ splits = look_up_table [node .identifier ]["graph_node" ].inputs [field_name ].split ("${{" )
120+ for split in splits [1 :]:
121+ if "}}" in split :
122+
123+ identifier = None
124+ field = None
125+
126+ syntax_string = split .split ("}}" )[0 ].strip ()
127+
128+ if syntax_string .startswith ("identifier." ) and len (syntax_string .split ("." )) == 3 :
129+ identifier = syntax_string .split ("." )[1 ].strip ()
130+ field = syntax_string .split ("." )[2 ].strip ()
131+ else :
132+ errors .append (f"{ node .node_name } .Inputs field '{ field_name } ' references field { syntax_string } which is not a valid output field" )
133+ continue
134+
135+ if identifier is None or field is None :
136+ errors .append (f"{ node .node_name } .Inputs field '{ field_name } ' references field { syntax_string } which is not a valid output field" )
137+ continue
138+
139+ if identifier not in dependencies_graph [node .identifier ]:
140+ errors .append (f"{ node .node_name } .Inputs field '{ field_name } ' references node { identifier } which is not a dependency of { node .identifier } " )
141+ continue
142+
143+ output_model_class = create_model (look_up_table [identifier ]["database_node" ].outputs_schema )
144+ if field not in output_model_class .model_fields .keys ():
145+ errors .append (f"{ node .node_name } .Inputs field '{ field_name } ' references field { field } of node { identifier } which is not a valid output field" )
146+ continue
147+
148+ except Exception as e :
149+ errors .append (f"Error creating input model for node { node .identifier } : { str (e )} " )
150+
151+ async def build_dependencies_graph (graph_nodes : list [NodeTemplate ]):
152+ dependency_graph = {}
153+ for node in graph_nodes :
154+ dependency_graph [node .identifier ] = set ()
155+ if node .next_nodes is None :
156+ continue
157+ for next_node in node .next_nodes :
158+ dependency_graph [next_node ].add (node .identifier )
159+ dependency_graph [next_node ] = dependency_graph [next_node ] | dependency_graph [node .identifier ]
160+ return dependency_graph
161+
162+ async def verify_topology (graph_nodes : list [NodeTemplate ], errors : list [str ]):
163+ # verify that the graph is a tree
164+ # verify that the graph is connected
165+ dependencies = {}
166+ identifier_to_node = {}
167+ visited = {}
168+
169+ for node in graph_nodes :
170+ if node .identifier in dependencies .keys ():
171+ errors .append (f"Multiple identifier { node .identifier } incorrect topology" )
172+ return
173+ dependencies [node .identifier ] = set ()
174+ identifier_to_node [node .identifier ] = node
175+ visited [node .identifier ] = False
176+
177+ # verify that there exists only one root node
178+ for node in graph_nodes :
179+ if node .next_nodes is None :
180+ continue
181+ for next_node in node .next_nodes :
182+ dependencies [next_node ].add (node .identifier )
183+
184+ # verify that there exists only one root node
185+ root_nodes = [node for node in graph_nodes if len (dependencies [node .identifier ]) == 0 ]
186+ if len (root_nodes ) != 1 :
187+ errors .append (f"Graph has { len (root_nodes )} root nodes, expected 1" )
188+ return
189+
190+ # verify that the graph is a tree
191+ to_visit = deque ([root_nodes [0 ].identifier ])
192+
193+ while len (to_visit ) > 0 :
194+ current_node = to_visit .popleft ()
195+ visited [current_node ] = True
196+
197+ if identifier_to_node [current_node ].next_nodes is None :
198+ continue
199+
200+ for next_node in identifier_to_node [current_node ].next_nodes :
201+ if visited [next_node ]:
202+ errors .append (f"Graph is not a tree at { current_node } -> { next_node } " )
203+ else :
204+ to_visit .append (next_node )
205+
206+ for identifier , visited_value in visited .items ():
207+ if not visited_value :
208+ errors .append (f"Graph is not connected at { identifier } " )
209+
71210async def verify_graph (graph_template : GraphTemplate ):
72211 try :
73212 errors = []
213+ database_nodes = await get_database_nodes (graph_template .nodes , graph_template .namespace )
214+
74215 await verify_nodes_names (graph_template .nodes , errors )
75216 await verify_nodes_namespace (graph_template .nodes , graph_template .namespace , errors )
76- await verify_node_exists (graph_template .nodes , graph_template . namespace , errors )
217+ await verify_node_exists (graph_template .nodes , database_nodes , errors )
77218 await verify_node_identifiers (graph_template .nodes , errors )
219+ await verify_secrets (graph_template , database_nodes , errors )
220+ await verify_topology (graph_template .nodes , errors )
78221
79222 if errors :
80223 graph_template .validation_status = GraphTemplateValidationStatus .INVALID
81224 graph_template .validation_errors = errors
82225 await graph_template .save ()
83226 return
84227
228+ dependencies_graph = await build_dependencies_graph (graph_template .nodes )
229+ await verify_inputs (graph_template .nodes , database_nodes , dependencies_graph , errors )
230+
85231 graph_template .validation_status = GraphTemplateValidationStatus .VALID
86232 graph_template .validation_errors = None
87233 await graph_template .save ()
0 commit comments