Skip to content

Commit 40218db

Browse files
committed
Refactor GraphTemplate model to enhance validation and structure
- Reintroduced validation methods for nodes, ensuring namespace consistency and connectivity to the root node. - Added checks for acyclic structure and existence of unit identifiers within nodes. - Moved the get_node_by_identifier and get_parents_by_identifier methods to the end of the class for better organization. - Removed outdated validation methods to streamline the model. These changes improve the integrity and validation of the GraphTemplate model, ensuring robust graph structure management.
1 parent d61f3ab commit 40218db

1 file changed

Lines changed: 61 additions & 60 deletions

File tree

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

Lines changed: 61 additions & 60 deletions
Original file line numberDiff line numberDiff line change
@@ -19,6 +19,7 @@ class GraphTemplate(BaseDatabaseModel):
1919
validation_status: GraphTemplateValidationStatus = Field(..., description="Validation status of the graph")
2020
validation_errors: Optional[List[str]] = Field(None, description="Validation errors of the graph")
2121
secrets: Dict[str, str] = Field(default_factory=dict, description="Secrets of the graph")
22+
2223
_node_by_identifier: Dict[str, NodeTemplate] | None = PrivateAttr(default=None)
2324
_parents_by_identifier: Dict[str, set[str]] | None = PrivateAttr(default=None)
2425
_root_node: NodeTemplate | None = PrivateAttr(default=None)
@@ -86,21 +87,6 @@ def dfs(node_identifier: str, parents: set[str]) -> None:
8687

8788
except Exception as e:
8889
raise ValueError(f"Error building dependency graph: {e}")
89-
90-
def get_node_by_identifier(self, identifier: str) -> NodeTemplate | None:
91-
"""Get a node by its identifier using O(1) dictionary lookup."""
92-
if self._node_by_identifier is None:
93-
self._build_node_by_identifier()
94-
95-
assert self._node_by_identifier is not None
96-
return self._node_by_identifier.get(identifier)
97-
98-
def get_parents_by_identifier(self, identifier: str) -> set[str]:
99-
if self._parents_by_identifier is None:
100-
self._build_parents_by_identifier()
101-
102-
assert self._parents_by_identifier is not None
103-
return self._parents_by_identifier[identifier]
10490

10591
@field_validator('name')
10692
@classmethod
@@ -130,50 +116,6 @@ def validate_secrets(cls, v: Dict[str, str]) -> Dict[str, str]:
130116

131117
return v
132118

133-
@model_validator(mode='after')
134-
def validate_nodes(self) -> Self:
135-
for node in self.nodes:
136-
if node.namespace != self.namespace:
137-
raise ValueError(f"Node namespace {node.namespace} does not match graph namespace {self.namespace}")
138-
return self
139-
140-
@model_validator(mode='after')
141-
def validate_graph_is_connected(self) -> Self:
142-
errors = []
143-
root_node_identifier = self.get_root_node().identifier
144-
for node in self.nodes:
145-
if root_node_identifier not in self.get_parents_by_identifier(node.identifier):
146-
errors.append(f"Node {node.identifier} is not connected to the root node")
147-
if errors:
148-
raise ValueError("\n".join(errors))
149-
return self
150-
151-
@model_validator(mode='after')
152-
def validate_graph_is_acyclic(self) -> Self:
153-
errors = []
154-
for node in self.nodes:
155-
if node.identifier in self.get_parents_by_identifier(node.identifier):
156-
errors.append(f"Node {node.identifier} is not acyclic")
157-
if errors:
158-
raise ValueError("\n".join(errors))
159-
return self
160-
161-
@model_validator(mode='after')
162-
def verify_unites_identifiers_exist(self) -> Self:
163-
errors = []
164-
identifiers = set()
165-
for node in self.nodes:
166-
identifiers.add(node.identifier)
167-
for node in self.nodes:
168-
if node.unites is not None:
169-
if node.unites.identifier not in identifiers:
170-
errors.append(f"Node {node.identifier} has a unit {node.unites.identifier} that does not exist")
171-
if node.unites.identifier == node.identifier:
172-
errors.append(f"Node {node.identifier} has a unit {node.unites.identifier} that is the same as the node itself")
173-
if errors:
174-
raise ValueError("\n".join(errors))
175-
return self
176-
177119
@field_validator('nodes')
178120
@classmethod
179121
def validate_unique_identifiers(cls, v: List[NodeTemplate]) -> List[NodeTemplate]:
@@ -218,6 +160,50 @@ def _validate_secret_value(cls, secret_value: str) -> None:
218160
raise ValueError("Decoded value is too short to contain valid nonce")
219161
except Exception:
220162
raise ValueError("Value is not valid URL-safe base64 encoded")
163+
164+
@model_validator(mode='after')
165+
def validate_nodes(self) -> Self:
166+
for node in self.nodes:
167+
if node.namespace != self.namespace:
168+
raise ValueError(f"Node namespace {node.namespace} does not match graph namespace {self.namespace}")
169+
return self
170+
171+
@model_validator(mode='after')
172+
def validate_graph_is_connected(self) -> Self:
173+
errors = []
174+
root_node_identifier = self.get_root_node().identifier
175+
for node in self.nodes:
176+
if root_node_identifier not in self.get_parents_by_identifier(node.identifier):
177+
errors.append(f"Node {node.identifier} is not connected to the root node")
178+
if errors:
179+
raise ValueError("\n".join(errors))
180+
return self
181+
182+
@model_validator(mode='after')
183+
def validate_graph_is_acyclic(self) -> Self:
184+
errors = []
185+
for node in self.nodes:
186+
if node.identifier in self.get_parents_by_identifier(node.identifier):
187+
errors.append(f"Node {node.identifier} is not acyclic")
188+
if errors:
189+
raise ValueError("\n".join(errors))
190+
return self
191+
192+
@model_validator(mode='after')
193+
def verify_unites_identifiers_exist(self) -> Self:
194+
errors = []
195+
identifiers = set()
196+
for node in self.nodes:
197+
identifiers.add(node.identifier)
198+
for node in self.nodes:
199+
if node.unites is not None:
200+
if node.unites.identifier not in identifiers:
201+
errors.append(f"Node {node.identifier} has a unit {node.unites.identifier} that does not exist")
202+
if node.unites.identifier == node.identifier:
203+
errors.append(f"Node {node.identifier} has a unit {node.unites.identifier} that is the same as the node itself")
204+
if errors:
205+
raise ValueError("\n".join(errors))
206+
return self
221207

222208
def set_secrets(self, secrets: Dict[str, str]) -> "GraphTemplate":
223209
self.secrets = {secret_name: get_encrypter().encrypt(secret_value) for secret_name, secret_value in secrets.items()}
@@ -247,6 +233,21 @@ def get_root_node(self) -> NodeTemplate:
247233
def is_validating(self) -> bool:
248234
return self.validation_status in (GraphTemplateValidationStatus.ONGOING, GraphTemplateValidationStatus.PENDING)
249235

236+
def get_node_by_identifier(self, identifier: str) -> NodeTemplate | None:
237+
"""Get a node by its identifier using O(1) dictionary lookup."""
238+
if self._node_by_identifier is None:
239+
self._build_node_by_identifier()
240+
241+
assert self._node_by_identifier is not None
242+
return self._node_by_identifier.get(identifier)
243+
244+
def get_parents_by_identifier(self, identifier: str) -> set[str]:
245+
if self._parents_by_identifier is None:
246+
self._build_parents_by_identifier()
247+
248+
assert self._parents_by_identifier is not None
249+
return self._parents_by_identifier[identifier]
250+
250251
@staticmethod
251252
async def get(namespace: str, graph_name: str) -> "GraphTemplate":
252253
graph_template = await GraphTemplate.find_one(GraphTemplate.namespace == namespace, GraphTemplate.name == graph_name)
@@ -275,4 +276,4 @@ async def get_valid(namespace: str, graph_name: str, polling_interval: float = 1
275276
await asyncio.sleep(polling_interval)
276277
else:
277278
raise ValueError(f"Graph template is in a non-validating state: {graph_template.validation_status.value} for namespace: {namespace} and graph name: {graph_name}")
278-
raise ValueError(f"Graph template is not valid for namespace: {namespace} and graph name: {graph_name} after {timeout} seconds")
279+
raise ValueError(f"Graph template is not valid for namespace: {namespace} and graph name: {graph_name} after {timeout} seconds")

0 commit comments

Comments
 (0)