Skip to content

Commit 2a9901a

Browse files
committed
Refactor models and verification functions for improved validation and clarity
- Updated the DependentString model to utilize enumeration for order in dependent placeholders, enhancing readability. - Refactored validation methods in NodeTemplate and GraphTemplate to ensure non-empty inputs and unique identifiers. - Improved error handling in verification functions to check for string types in inputs and ensure proper validation of graph structures. - Streamlined the RegisteredNode model's list_nodes_by_templates method to handle empty templates gracefully. These changes enhance the robustness and maintainability of the models and their validation processes.
1 parent 01c4bfc commit 2a9901a

5 files changed

Lines changed: 53 additions & 45 deletions

File tree

state-manager/app/models/db/graph_template_model.py

Lines changed: 21 additions & 22 deletions
Original file line numberDiff line numberDiff line change
@@ -41,14 +41,12 @@ def _build_root_node(self) -> None:
4141
in_degree = {node.identifier: 0 for node in self.nodes}
4242

4343
for node in self.nodes:
44-
if node.next_nodes is None:
45-
continue
46-
for next_node in node.next_nodes:
47-
in_degree[next_node] += 1
44+
if node.next_nodes is not None:
45+
for next_node in node.next_nodes:
46+
in_degree[next_node] += 1
4847

49-
if node.unites is None:
50-
continue
51-
in_degree[node.identifier] += 1
48+
if node.unites is not None:
49+
in_degree[node.identifier] += 1
5250

5351
zero_in_degree_nodes = [node for node in self.nodes if in_degree[node.identifier] == 0]
5452
if len(zero_in_degree_nodes) != 1:
@@ -58,29 +56,31 @@ def _build_root_node(self) -> None:
5856
def _build_parents_by_identifier(self) -> None:
5957
try:
6058
root_node_identifier = self.get_root_node().identifier
61-
self._parents_by_identifier = {
62-
node.identifier: set() for node in self.nodes
63-
}
6459

65-
visited = set()
60+
visited = {}
61+
62+
self._parents_by_identifier = {}
63+
for node in self.nodes:
64+
self._parents_by_identifier[node.identifier] = set()
65+
visited[node.identifier] = False
6666

6767
def dfs(node_identifier: str, parents: set[str]) -> None:
6868
assert self._parents_by_identifier is not None
6969

7070
self._parents_by_identifier[node_identifier] = parents | self._parents_by_identifier[node_identifier]
7171

72-
if node_identifier in visited:
72+
if visited[node_identifier]:
7373
return
7474

75-
visited.add(node_identifier)
75+
visited[node_identifier] = True
7676

7777
node = self.get_node_by_identifier(node_identifier)
7878
if node is None:
7979
return
8080
if node.next_nodes is None:
8181
return
8282
if node.unites is not None:
83-
self._parents_by_identifier[node.unites.identifier].add(node_identifier)
83+
self._parents_by_identifier[node_identifier].add(node.unites.identifier)
8484
for next_node_identifier in node.next_nodes:
8585
dfs(next_node_identifier, parents | {node_identifier})
8686

@@ -162,18 +162,13 @@ def _validate_secret_value(cls, secret_value: str) -> None:
162162
except Exception:
163163
raise ValueError("Value is not valid URL-safe base64 encoded")
164164

165-
@model_validator(mode='after')
166-
def validate_nodes(self) -> Self:
167-
for node in self.nodes:
168-
if node.namespace != self.namespace:
169-
raise ValueError(f"Node namespace {node.namespace} does not match graph namespace {self.namespace}")
170-
return self
171-
172165
@model_validator(mode='after')
173166
def validate_graph_is_connected(self) -> Self:
174167
errors = []
175168
root_node_identifier = self.get_root_node().identifier
176169
for node in self.nodes:
170+
if node.identifier == root_node_identifier:
171+
continue
177172
if root_node_identifier not in self.get_parents_by_identifier(node.identifier):
178173
errors.append(f"Node {node.identifier} is not connected to the root node")
179174
if errors:
@@ -213,6 +208,10 @@ def verify_input_dependencies(self) -> Self:
213208
for node in self.nodes:
214209
for input_value in node.inputs.values():
215210
try:
211+
if not isinstance(input_value, str):
212+
errors.append(f"Input {input_value} is not a string")
213+
continue
214+
216215
dependent_string = DependentString.create_dependent_string(input_value)
217216
dependent_identifiers = set([identifier for identifier, _ in dependent_string.get_identifier_field()])
218217

@@ -272,7 +271,7 @@ def get_parents_by_identifier(self, identifier: str) -> set[str]:
272271
self._build_parents_by_identifier()
273272

274273
assert self._parents_by_identifier is not None
275-
return self._parents_by_identifier[identifier]
274+
return self._parents_by_identifier.get(identifier, set())
276275

277276
@staticmethod
278277
async def get(namespace: str, graph_name: str) -> "GraphTemplate":

state-manager/app/models/db/registered_node.py

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -32,11 +32,13 @@ async def get_by_name_and_namespace(name: str, namespace: str) -> "RegisteredNod
3232

3333
@staticmethod
3434
async def list_nodes_by_templates(templates: list[NodeTemplate]) -> list["RegisteredNode"]:
35+
if len(templates) == 0:
36+
return []
37+
3538
query = {
3639
"$or": [
3740
{"name": node.node_name, "namespace": node.namespace}
3841
for node in templates
3942
]
4043
}
41-
nodes = await RegisteredNode.find(query).to_list()
42-
return [node for node in nodes if node is not None]
44+
return await RegisteredNode.find(query).to_list()

state-manager/app/models/dependent_string.py

Lines changed: 2 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -27,19 +27,17 @@ def create_dependent_string(syntax_string: str) -> "DependentString":
2727
return DependentString(head=syntax_string, dependents={})
2828

2929
dependent_string = DependentString(head=splits[0], dependents={})
30-
order = 0
3130

32-
for split in splits[1:]:
31+
for order, split in enumerate(splits[1:]):
3332
if "}}" not in split:
3433
raise ValueError(f"Invalid syntax string placeholder {split} for: {syntax_string} '${{' not closed")
35-
placeholder_content, tail = split.split("}}")
34+
placeholder_content, tail = split.split("}}", 1)
3635

3736
parts = [p.strip() for p in placeholder_content.split(".")]
3837
if len(parts) != 3 or parts[1] != "outputs":
3938
raise ValueError(f"Invalid syntax string placeholder {placeholder_content} for: {syntax_string}")
4039

4140
dependent_string.dependents[order] = Dependent(identifier=parts[0], field=parts[2], tail=tail)
42-
order += 1
4341

4442
return dependent_string
4543

state-manager/app/models/node_template_model.py

Lines changed: 9 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -22,16 +22,9 @@ def validate_node_name(cls, v: str) -> str:
2222
raise ValueError("Node name cannot be empty")
2323
return v
2424

25-
@field_validator('node_name')
26-
@classmethod
27-
def validate_node_name_unique(cls, v: str) -> str:
28-
if v == "" or v is None:
29-
raise ValueError("Node name cannot be empty")
30-
return v
31-
3225
@field_validator('identifier')
3326
@classmethod
34-
def validate_identifier_unique(cls, v: str) -> str:
27+
def validate_identifier(cls, v: str) -> str:
3528
if v == "" or v is None:
3629
raise ValueError("Node identifier cannot be empty")
3730
return v
@@ -43,10 +36,15 @@ def validate_next_nodes(cls, v: Optional[List[str]]) -> Optional[List[str]]:
4336
errors = []
4437
if v is not None:
4538
for next_node_identifier in v:
39+
4640
if next_node_identifier == "" or next_node_identifier is None:
4741
errors.append("Next node identifier cannot be empty")
48-
elif next_node_identifier in identifiers:
42+
continue
43+
44+
if next_node_identifier in identifiers:
4945
errors.append(f"Next node identifier {next_node_identifier} is not unique")
46+
continue
47+
5048
identifiers.add(next_node_identifier)
5149
if errors:
5250
raise ValueError("\n".join(errors))
@@ -63,5 +61,7 @@ def validate_unites(cls, v: Optional[Unites]) -> Optional[Unites]:
6361
def get_dependent_strings(self) -> list[DependentString]:
6462
dependent_strings = []
6563
for input_value in self.inputs.values():
64+
if not isinstance(input_value, str):
65+
raise ValueError(f"Input {input_value} is not a string")
6666
dependent_strings.append(DependentString.create_dependent_string(input_value))
6767
return dependent_strings

state-manager/app/tasks/verify_graph.py

Lines changed: 17 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -43,22 +43,22 @@ async def verify_secrets(graph_template: GraphTemplate, registered_nodes: list[R
4343
async def verify_inputs(graph_template: GraphTemplate, registered_nodes: list[RegisteredNode]) -> list[str]:
4444
errors = []
4545
look_up_table = {
46-
(node.node_name, node.namespace): node
47-
for node in registered_nodes
46+
(rn.node_name, rn.namespace): rn
47+
for rn in registered_nodes
4848
}
4949

5050
for node in graph_template.nodes:
5151
if node.inputs is None:
5252
continue
5353

54-
if (node.node_name, node.namespace) not in look_up_table:
54+
registered_node = look_up_table.get((node.node_name, node.namespace))
55+
if registered_node is None:
5556
errors.append(f"Node {node.node_name} in namespace {node.namespace} does not exist")
5657
continue
5758

58-
registered_node = look_up_table[(node.node_name, node.namespace)]
59-
registerd_node_input_model = create_model(registered_node.inputs_schema)
59+
registered_node_input_model = create_model(registered_node.inputs_schema)
6060

61-
for input_name, input_info in registerd_node_input_model.model_fields.items():
61+
for input_name, input_info in registered_node_input_model.model_fields.items():
6262
if input_info.annotation is not str:
6363
errors.append(f"Input {input_name} in node {node.node_name} in namespace {node.namespace} is not a string")
6464
continue
@@ -75,7 +75,7 @@ async def verify_inputs(graph_template: GraphTemplate, registered_nodes: list[Re
7575
temp_node = graph_template.get_node_by_identifier(identifier)
7676
assert temp_node is not None
7777

78-
registered_node = look_up_table[(temp_node.node_name, temp_node.namespace)]
78+
registered_node = look_up_table.get((temp_node.node_name, temp_node.namespace))
7979
if registered_node is None:
8080
errors.append(f"Node {temp_node.node_name} in namespace {temp_node.namespace} does not exist")
8181
continue
@@ -100,7 +100,16 @@ async def verify_graph(graph_template: GraphTemplate):
100100
verify_secrets(graph_template, registered_nodes),
101101
verify_inputs(graph_template, registered_nodes)
102102
]
103-
errors.extend(await asyncio.gather(*basic_verify_tasks))
103+
resultant_errors = await asyncio.gather(*basic_verify_tasks)
104+
105+
for error in resultant_errors:
106+
errors.extend(error)
107+
108+
if len(errors) > 0:
109+
graph_template.validation_status = GraphTemplateValidationStatus.INVALID
110+
graph_template.validation_errors = errors
111+
await graph_template.save()
112+
return
104113

105114
graph_template.validation_status = GraphTemplateValidationStatus.VALID
106115
graph_template.validation_errors = None

0 commit comments

Comments
 (0)