Skip to content

Commit b3ee4b9

Browse files
committed
Add DependentString model and integrate into NodeTemplate and GraphTemplate
- Introduced a new DependentString model to manage dependent values and their relationships. - Updated NodeTemplate to include a method for generating dependent strings from input values. - Enhanced GraphTemplate with validation for input dependencies using the new DependentString model. - Refactored existing code to utilize the new model, improving clarity and maintainability. These changes enhance the handling of dependent values within templates, ensuring better validation and structure in the graph management process.
1 parent 8377710 commit b3ee4b9

4 files changed

Lines changed: 110 additions & 53 deletions

File tree

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

Lines changed: 26 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,7 @@
1010
from ..graph_template_validation_status import GraphTemplateValidationStatus
1111
from ..node_template_model import NodeTemplate
1212
from app.utils.encrypter import get_encrypter
13+
from app.models.dependent_string import DependentString
1314

1415

1516
class GraphTemplate(BaseDatabaseModel):
@@ -204,6 +205,31 @@ def verify_unites_identifiers_exist(self) -> Self:
204205
if errors:
205206
raise ValueError("\n".join(errors))
206207
return self
208+
209+
@model_validator(mode='after')
210+
def verify_input_dependencies(self) -> Self:
211+
errors = []
212+
213+
for node in self.nodes:
214+
for input_value in node.inputs.values():
215+
try:
216+
dependent_string = DependentString.create_dependent_string(input_value)
217+
dependent_identifiers = set([identifier for identifier, _ in dependent_string.get_identifier_field()])
218+
219+
for identifier in dependent_identifiers:
220+
if node.unites is not None:
221+
if identifier not in self.get_parents_by_identifier(node.unites.identifier):
222+
errors.append(f"Input {input_value} depends on {identifier} but {identifier} is not a parent of unites {node.unites.identifier}")
223+
else:
224+
if identifier not in self.get_parents_by_identifier(node.identifier):
225+
errors.append(f"Input {input_value} depends on {identifier} but {identifier} is not a parent of {node.identifier}")
226+
227+
except Exception as e:
228+
errors.append(f"Error creating dependent string for input {input_value}: {e}")
229+
if errors:
230+
raise ValueError("\n".join(errors))
231+
232+
return self
207233

208234
def set_secrets(self, secrets: Dict[str, str]) -> "GraphTemplate":
209235
self.secrets = {secret_name: get_encrypter().encrypt(secret_value) for secret_name, secret_value in secrets.items()}
Lines changed: 64 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,64 @@
1+
from pydantic import BaseModel, PrivateAttr
2+
3+
class Dependent(BaseModel):
4+
identifier: str
5+
field: str
6+
tail: str
7+
value: str | None = None
8+
9+
class DependentString(BaseModel):
10+
head: str
11+
dependents: dict[int, Dependent]
12+
_mapping_key_to_dependent: dict[tuple[str, str], list[Dependent]] = PrivateAttr(default_factory=dict)
13+
14+
def generate_string(self) -> str:
15+
base = self.head
16+
for key in sorted(self.dependents.keys()):
17+
dependent = self.dependents[key]
18+
if dependent.value is None:
19+
raise ValueError(f"Dependent value is not set for: {dependent}")
20+
base += dependent.value + dependent.tail
21+
return base
22+
23+
@staticmethod
24+
def create_dependent_string(syntax_string: str) -> "DependentString":
25+
splits = syntax_string.split("${{")
26+
if len(splits) <= 1:
27+
return DependentString(head=syntax_string, dependents={})
28+
29+
dependent_string = DependentString(head=splits[0], dependents={})
30+
order = 0
31+
32+
for split in splits[1:]:
33+
if "}}" not in split:
34+
raise ValueError(f"Invalid syntax string placeholder {split} for: {syntax_string} '${{' not closed")
35+
placeholder_content, tail = split.split("}}")
36+
37+
parts = [p.strip() for p in placeholder_content.split(".")]
38+
if len(parts) != 3 or parts[1] != "outputs":
39+
raise ValueError(f"Invalid syntax string placeholder {placeholder_content} for: {syntax_string}")
40+
41+
dependent_string.dependents[order] = Dependent(identifier=parts[0], field=parts[2], tail=tail)
42+
order += 1
43+
44+
return dependent_string
45+
46+
def _build_mapping_key_to_dependent(self):
47+
if self._mapping_key_to_dependent is not None:
48+
return
49+
50+
for dependent in self.dependents.values():
51+
mapping_key = (dependent.identifier, dependent.field)
52+
if mapping_key not in self._mapping_key_to_dependent:
53+
self._mapping_key_to_dependent[mapping_key] = []
54+
self._mapping_key_to_dependent[mapping_key].append(dependent)
55+
56+
def set_value(self, identifier: str, field: str, value: str):
57+
self._build_mapping_key_to_dependent()
58+
mapping_key = (identifier, field)
59+
for dependent in self._mapping_key_to_dependent[mapping_key]:
60+
dependent.value = value
61+
62+
def get_identifier_field(self) -> list[tuple[str, str]]:
63+
self._build_mapping_key_to_dependent()
64+
return list(self._mapping_key_to_dependent.keys())

state-manager/app/models/node_template_model.py

Lines changed: 8 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,6 @@
11
from pydantic import Field, BaseModel, field_validator
22
from typing import Any, Optional, List
3+
from .dependent_string import DependentString
34

45

56
class Unites(BaseModel):
@@ -57,4 +58,10 @@ def validate_unites(cls, v: Optional[Unites]) -> Optional[Unites]:
5758
if v is not None:
5859
if v.identifier == "" or v.identifier is None:
5960
raise ValueError("Unites identifier cannot be empty")
60-
return v
61+
return v
62+
63+
def get_dependent_strings(self) -> list[DependentString]:
64+
dependent_strings = []
65+
for input_value in self.inputs.values():
66+
dependent_strings.append(DependentString.create_dependent_string(input_value))
67+
return dependent_strings

state-manager/app/tasks/create_next_states.py

Lines changed: 12 additions & 52 deletions
Original file line numberDiff line numberDiff line change
@@ -7,31 +7,13 @@
77
from app.models.state_status_enum import StateStatusEnum
88
from app.models.node_template_model import NodeTemplate
99
from app.models.db.registered_node import RegisteredNode
10+
from app.models.dependent_string import DependentString
1011
from json_schema_to_pydantic import create_model
1112
from pydantic import BaseModel
1213
from typing import Type
1314

1415
logger = LogsManager().get_logger()
1516

16-
class Dependent(BaseModel):
17-
identifier: str
18-
field: str
19-
tail: str
20-
value: str | None = None
21-
22-
class DependentString(BaseModel):
23-
head: str
24-
dependents: dict[int, Dependent]
25-
26-
def generate_string(self) -> str:
27-
base = self.head
28-
for key in sorted(self.dependents.keys()):
29-
dependent = self.dependents[key]
30-
if dependent.value is None:
31-
raise ValueError(f"Dependent value is not set for: {dependent}")
32-
base += dependent.value + dependent.tail
33-
return base
34-
3517
async def mark_success_states(state_ids: list[PydanticObjectId]):
3618
await State.find(
3719
In(State.id, state_ids)
@@ -60,36 +42,14 @@ async def check_unites_satisfied(namespace: str, graph_name: str, node_template:
6042
return True
6143

6244

63-
def get_dependents(syntax_string: str) -> DependentString:
64-
splits = syntax_string.split("${{")
65-
if len(splits) <= 1:
66-
return DependentString(head=syntax_string, dependents={})
67-
68-
dependent_string = DependentString(head=splits[0], dependents={})
69-
order = 0
70-
71-
for split in splits[1:]:
72-
if "}}" not in split:
73-
raise ValueError(f"Invalid syntax string placeholder {split} for: {syntax_string} '${{' not closed")
74-
placeholder_content, tail = split.split("}}")
75-
76-
parts = [p.strip() for p in placeholder_content.split(".")]
77-
if len(parts) != 3 or parts[1] != "outputs":
78-
raise ValueError(f"Invalid syntax string placeholder {placeholder_content} for: {syntax_string}")
79-
80-
dependent_string.dependents[order] = Dependent(identifier=parts[0], field=parts[2], tail=tail)
81-
order += 1
82-
83-
return dependent_string
84-
8545
def validate_dependencies(next_state_node_template: NodeTemplate, next_state_input_model: Type[BaseModel], identifier: str, parents: dict[str, State]) -> None:
8646
"""Validate that all dependencies exist before processing them."""
8747
# 1) Confirm each model field is present in next_state_node_template.inputs
8848
for field_name in next_state_input_model.model_fields.keys():
8949
if field_name not in next_state_node_template.inputs:
9050
raise ValueError(f"Field '{field_name}' not found in inputs for template '{next_state_node_template.identifier}'")
9151

92-
dependency_string = get_dependents(next_state_node_template.inputs[field_name])
52+
dependency_string = DependentString.create_dependent_string(next_state_node_template.inputs[field_name])
9353

9454
for dependent in dependency_string.dependents.values():
9555
# 2) For each placeholder, verify the identifier is either current or present in parents
@@ -110,16 +70,16 @@ def generate_next_state(next_state_input_model: Type[BaseModel], next_state_node
11070
next_state_input_data = {}
11171

11272
for field_name, _ in next_state_input_model.model_fields.items():
113-
dependency_string = get_dependents(next_state_node_template.inputs[field_name])
114-
115-
for key in sorted(dependency_string.dependents.keys()):
116-
if dependency_string.dependents[key].identifier == current_state.identifier:
117-
if dependency_string.dependents[key].field not in current_state.outputs:
118-
raise AttributeError(f"Output field '{dependency_string.dependents[key].field}' not found on current state '{current_state.identifier}' for template '{next_state_node_template.identifier}'")
119-
dependency_string.dependents[key].value = current_state.outputs[dependency_string.dependents[key].field]
120-
else:
121-
dependency_string.dependents[key].value = parents[dependency_string.dependents[key].identifier].outputs[dependency_string.dependents[key].field]
122-
73+
dependency_string = DependentString.create_dependent_string(next_state_node_template.inputs[field_name])
74+
75+
for identifier, field in dependency_string.get_identifier_field():
76+
if identifier == current_state.identifier:
77+
if field not in current_state.outputs:
78+
raise AttributeError(f"Output field '{field}' not found on current state '{current_state.identifier}' for template '{next_state_node_template.identifier}'")
79+
dependency_string.set_value(identifier, field, current_state.outputs[field])
80+
else:
81+
dependency_string.set_value(identifier, field, parents[identifier].outputs[field])
82+
12383
next_state_input_data[field_name] = dependency_string.generate_string()
12484

12585
new_parents = {

0 commit comments

Comments
 (0)