stageflow-framework 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,71 @@
1
+ from dataclasses import dataclass
2
+ from datetime import datetime, UTC
3
+
4
+
5
+ @dataclass
6
+ class InputSpec:
7
+ type: str | None = None
8
+ description: str | None = None
9
+ required: bool = True
10
+ default: str | int | float | bool | None = None
11
+ payload_schema: object | None = None
12
+
13
+
14
+ @dataclass
15
+ class EventSpec:
16
+ type: str
17
+ description: str | None = None
18
+ payload_schema: object | None = None
19
+
20
+
21
+ class Event:
22
+ type: str
23
+ session_id: str
24
+ node_id: str | None
25
+ stage_id: str | None
26
+ action_id: str | None
27
+ payload: dict
28
+ timestamp: datetime
29
+
30
+ def __init__(
31
+ self,
32
+ type: str,
33
+ session_id: str,
34
+ node_id: str | None = None,
35
+ stage_id: str | None = None,
36
+ action_id: str | None = None,
37
+ payload: dict = None,
38
+ ):
39
+ self.type = type
40
+ self.session_id = session_id
41
+ self.node_id = node_id
42
+ self.stage_id = stage_id
43
+ self.action_id = action_id
44
+ self.payload = payload or {}
45
+ self.timestamp = datetime.now(UTC)
46
+
47
+ def to_dict(self) -> dict:
48
+ return {
49
+ "type": self.type,
50
+ "session_id": self.session_id,
51
+ "node_id": self.node_id,
52
+ "stage_id": self.stage_id,
53
+ "action_id": self.action_id,
54
+ "payload": self.payload,
55
+ "timestamp": self.timestamp.isoformat(),
56
+ }
57
+
58
+ @classmethod
59
+ def from_dict(cls, data: dict) -> "Event":
60
+ ts_raw = data.get("timestamp")
61
+ timestamp = datetime.fromisoformat(ts_raw) if ts_raw else datetime.now(UTC)
62
+ ev = cls(
63
+ type=data.get("type"),
64
+ session_id=data.get("session_id"),
65
+ node_id=data.get("node_id"),
66
+ stage_id=data.get("stage_id"),
67
+ action_id=data.get("action_id"),
68
+ payload=data.get("payload"),
69
+ )
70
+ ev.timestamp = timestamp
71
+ return ev
@@ -0,0 +1,43 @@
1
+ from .context import Context
2
+
3
+
4
+ class JsonLogic:
5
+ def __init__(self, condition: dict):
6
+ self.condition = condition
7
+
8
+ def evaluate(self, context: Context) -> bool:
9
+ base = {"payload": context.payload}
10
+ if isinstance(context.payload, dict):
11
+ base.update(context.payload)
12
+ return self._eval(self.condition, base)
13
+
14
+ def _eval(self, expr, data):
15
+ if isinstance(expr, dict):
16
+ if len(expr) != 1:
17
+ raise ValueError("Invalid JSONLogic expression")
18
+ op, args = next(iter(expr.items()))
19
+ if not isinstance(args, list):
20
+ args = [args]
21
+
22
+ if op == "var":
23
+ path = args[0].split(".")
24
+ val = data
25
+ for p in path:
26
+ if val is None:
27
+ return None
28
+ val = val.get(p, None)
29
+ return val
30
+ elif op == "<":
31
+ return self._eval(args[0], data) < self._eval(args[1], data)
32
+ elif op == ">":
33
+ return self._eval(args[0], data) > self._eval(args[1], data)
34
+ elif op == "==":
35
+ return self._eval(args[0], data) == self._eval(args[1], data)
36
+ elif op == "and":
37
+ return all(self._eval(a, data) for a in args)
38
+ elif op == "or":
39
+ return any(self._eval(a, data) for a in args)
40
+ else:
41
+ raise NotImplementedError(f"Operator {op} not implemented")
42
+ else:
43
+ return expr
stageflow/core/node.py ADDED
@@ -0,0 +1,256 @@
1
+ from dataclasses import dataclass
2
+ from typing import Literal
3
+
4
+ from .jsonlogic import JsonLogic
5
+ from .stage import get_stage
6
+
7
+
8
+ class Node:
9
+ id: str
10
+ type: str
11
+ metadata: dict
12
+
13
+ def __init__(self, id: str, type: str, metadata: dict = None):
14
+ self.id = id
15
+ self.type = type
16
+ self.metadata = metadata or {}
17
+
18
+ @staticmethod
19
+ def from_dict(data: dict) -> "Node":
20
+ node_type = data.get("type")
21
+ match node_type:
22
+ case "condition":
23
+ return ConditionNode.from_dict(data)
24
+ case "parallel":
25
+ return ParallelNode.from_dict(data)
26
+ case "terminal":
27
+ return TerminalNode.from_dict(data)
28
+ case "stage":
29
+ return StageNode.from_dict(data)
30
+ case "subpipeline":
31
+ return SubPipelineNode.from_dict(data)
32
+ case "map":
33
+ return MapNode.from_dict(data)
34
+ case _:
35
+ raise ValueError(f"Unknown node type: {node_type}")
36
+
37
+
38
+ @dataclass
39
+ class Condition:
40
+ if_condition: JsonLogic
41
+ then_goto: str
42
+
43
+
44
+ class ConditionNode(Node):
45
+ conditions: list[Condition]
46
+ else_goto: str | None
47
+
48
+ def __init__(
49
+ self,
50
+ id: str,
51
+ type: str,
52
+ conditions: list[Condition],
53
+ else_goto: str | None = None,
54
+ metadata: dict = None,
55
+ ):
56
+ super().__init__(id, type, metadata)
57
+ self.conditions = conditions
58
+ self.else_goto = else_goto
59
+
60
+ @staticmethod
61
+ def from_dict(data: dict) -> "ConditionNode":
62
+ id = data.get("id")
63
+ type = data.get("type")
64
+ conditions_data = data.get("conditions", [])
65
+ conditions = [
66
+ Condition(if_condition=JsonLogic(cond["if"]), then_goto=cond["then"])
67
+ for cond in conditions_data
68
+ ]
69
+ else_goto = data.get("else")
70
+ metadata = data.get("metadata", {})
71
+ return ConditionNode(
72
+ id=id,
73
+ type=type,
74
+ conditions=conditions,
75
+ else_goto=else_goto,
76
+ metadata=metadata,
77
+ )
78
+
79
+
80
+ class ParallelNode(Node):
81
+ children: list[str]
82
+ policy: Literal["all", "any"]
83
+ next: str | None
84
+ cancel_on_error: bool
85
+
86
+ def __init__(
87
+ self,
88
+ id: str,
89
+ type: str,
90
+ children: list[str],
91
+ policy: Literal["all", "any"] = "all",
92
+ next: str | None = None,
93
+ cancel_on_error: bool = True,
94
+ metadata: dict = None,
95
+ ):
96
+ super().__init__(id, type, metadata)
97
+ self.children = children
98
+ self.policy = policy
99
+ self.next = next
100
+ self.cancel_on_error = cancel_on_error
101
+
102
+ @staticmethod
103
+ def from_dict(data: dict) -> "ParallelNode":
104
+ id = data.get("id")
105
+ type = data.get("type")
106
+ children = data.get("children", [])
107
+ policy = data.get("policy", "all")
108
+ next = data.get("next")
109
+ cancel_on_error = data.get("cancel_on_error", True)
110
+ metadata = data.get("metadata", {})
111
+ return ParallelNode(
112
+ id=id,
113
+ type=type,
114
+ children=children,
115
+ policy=policy,
116
+ next=next,
117
+ cancel_on_error=cancel_on_error,
118
+ metadata=metadata,
119
+ )
120
+
121
+
122
+ class TerminalNode(Node):
123
+ artifact_paths: list[str]
124
+ result: dict | None
125
+
126
+ def __init__(
127
+ self,
128
+ id: str,
129
+ type: str,
130
+ artifact_paths: list[str] = None,
131
+ result: dict | None = None,
132
+ metadata: dict = None,
133
+ ):
134
+ super().__init__(id, type, metadata)
135
+ self.artifact_paths = artifact_paths or []
136
+ self.result = result
137
+
138
+ @staticmethod
139
+ def from_dict(data: dict) -> "TerminalNode":
140
+ id = data.get("id")
141
+ type = data.get("type")
142
+ artifact_paths = data.get("artifacts", [])
143
+ result = data.get("result")
144
+ metadata = data.get("metadata", {})
145
+ return TerminalNode(
146
+ id=id,
147
+ type=type,
148
+ artifact_paths=artifact_paths,
149
+ result=result,
150
+ metadata=metadata,
151
+ )
152
+
153
+
154
+ class SubPipelineNode(Node):
155
+ subpipeline_id: str
156
+ inputs: dict
157
+ artifact_outputs: dict
158
+ result_output: str | None
159
+ next: str | None
160
+
161
+ def __init__(
162
+ self,
163
+ id: str,
164
+ type: str,
165
+ subpipeline_id: str,
166
+ inputs: dict = None,
167
+ artifact_outputs: dict = None,
168
+ result_output: str | None = None,
169
+ next: str | None = None,
170
+ metadata: dict = None,
171
+ ):
172
+ super().__init__(id, type, metadata)
173
+ self.subpipeline_id = subpipeline_id
174
+ self.inputs = inputs or {}
175
+ self.artifact_outputs = artifact_outputs or {}
176
+ self.result_output = result_output
177
+ self.next = next
178
+
179
+ @staticmethod
180
+ def from_dict(data: dict) -> "SubPipelineNode":
181
+ id = data.get("id")
182
+ type = data.get("type")
183
+ subpipeline_id = data.get("subpipeline_id")
184
+ if not subpipeline_id:
185
+ raise ValueError("subpipeline_id is required for subpipeline node")
186
+ inputs = data.get("inputs", {})
187
+ artifact_outputs = data.get("artifact_outputs", {})
188
+ result_output = data.get("result_output")
189
+ next = data.get("next")
190
+ metadata = data.get("metadata", {})
191
+ return SubPipelineNode(
192
+ id=id,
193
+ type=type,
194
+ subpipeline_id=subpipeline_id,
195
+ inputs=inputs,
196
+ artifact_outputs=artifact_outputs,
197
+ result_output=result_output,
198
+ next=next,
199
+ metadata=metadata,
200
+ )
201
+
202
+ class StageNode(Node):
203
+ stage: str
204
+ config: dict = {}
205
+ inputs: dict = {}
206
+ outputs: dict = {}
207
+ next: str | None = None
208
+ fallback: str | None = None
209
+
210
+ def __init__(
211
+ self,
212
+ id: str,
213
+ type: str,
214
+ stage: str,
215
+ config: dict = None,
216
+ arguments: dict = None,
217
+ outputs: dict = None,
218
+ next: str | None = None,
219
+ fallback: str | None = None,
220
+ metadata: dict = None,
221
+ ):
222
+ super().__init__(id, type, metadata)
223
+ if not get_stage(stage):
224
+ raise ValueError(f"Stage '{stage}' not found in registry")
225
+ self.stage = stage
226
+ self.config = config or {}
227
+ self.arguments = arguments or {}
228
+ self.outputs = outputs or {}
229
+ self.next = next
230
+ self.fallback = fallback
231
+
232
+ @staticmethod
233
+ def from_dict(data: dict) -> "StageNode":
234
+ id = data.get("id")
235
+ type = data.get("type")
236
+ stage = data.get("stage")
237
+ config = data.get("config", {})
238
+ arguments = data.get("arguments", {})
239
+ outputs = data.get("outputs", {})
240
+ next = data.get("next")
241
+ fallback = data.get("fallback")
242
+ metadata = data.get("metadata", {})
243
+ return StageNode(
244
+ id=id,
245
+ type=type,
246
+ stage=stage,
247
+ config=config,
248
+ arguments=arguments,
249
+ outputs=outputs,
250
+ next=next,
251
+ fallback=fallback,
252
+ metadata=metadata,
253
+ )
254
+
255
+ def get_stage_class(self):
256
+ return get_stage(self.stage)
@@ -0,0 +1,108 @@
1
+ from .node import Node, StageNode, ConditionNode, ParallelNode, TerminalNode, SubPipelineNode, MapNode
2
+ from stageflow.docs.schema import load_pipeline_schema
3
+
4
+ try:
5
+ from jsonschema import validate, ValidationError
6
+ except ImportError: # pragma: no cover - optional dependency
7
+ validate = None
8
+
9
+ class ValidationError(Exception):
10
+ pass
11
+
12
+
13
+ class Pipeline:
14
+ entry: str
15
+ nodes: list[Node]
16
+ _nodes_map: dict[str, Node]
17
+ metadata: dict
18
+ raw_json: dict
19
+ subpipelines: dict
20
+
21
+ def __init__(
22
+ self,
23
+ entry: str,
24
+ nodes: list[Node],
25
+ metadata: dict = None,
26
+ raw_json: dict = None,
27
+ subpipelines: dict | None = None,
28
+ ):
29
+ self.entry = entry
30
+ self.nodes = nodes
31
+ self._nodes_map = {node.id: node for node in nodes}
32
+ self.metadata = metadata or {}
33
+ self.raw_json = raw_json or {}
34
+ self.subpipelines = subpipelines or {}
35
+
36
+ @staticmethod
37
+ def from_dict(data: dict) -> "Pipeline":
38
+ # Schema validation before constructing nodes (if jsonschema available).
39
+ schema = None
40
+ if validate:
41
+ try:
42
+ schema = load_pipeline_schema()
43
+ validate(instance=data, schema=schema)
44
+ except ValidationError as e:
45
+ raise ValueError(f"Pipeline schema validation failed: {e.message}") from e
46
+ entry = data.get("entry")
47
+ nodes_data = data.get("nodes", [])
48
+ nodes = [Node.from_dict(node) for node in nodes_data]
49
+ metadata = data.get("metadata", {})
50
+ subpipelines = data.get("subpipelines", {})
51
+ # Validate subpipelines recursively.
52
+ if validate and schema:
53
+ for sub_id, sub_data in subpipelines.items():
54
+ try:
55
+ validate(instance=sub_data, schema=schema)
56
+ except ValidationError as e:
57
+ raise ValueError(f"Subpipeline '{sub_id}' schema validation failed: {e.message}") from e
58
+ return Pipeline(entry=entry, nodes=nodes, metadata=metadata, raw_json=data, subpipelines=subpipelines)
59
+
60
+ def get_node(self, node_id: str) -> Node:
61
+ if node_id not in self._nodes_map:
62
+ raise ValueError(f"Node with id {node_id} not found in pipeline")
63
+ return self._nodes_map[node_id]
64
+
65
+ def get_entry_node(self) -> Node:
66
+ return self.get_node(self.entry)
67
+
68
+ def validate(self) -> bool:
69
+ if self.entry not in self._nodes_map:
70
+ raise ValueError(f"Entry node '{self.entry}' not found in pipeline")
71
+
72
+ for node in self.nodes:
73
+ match node:
74
+ case StageNode():
75
+ stage_class = node.get_stage_class()
76
+ if not stage_class:
77
+ raise ValueError(f"Stage class for node {node.id} not found")
78
+ if node.next and node.next not in self._nodes_map:
79
+ raise ValueError(f"Next node '{node.next}' for stage {node.id} not found in pipeline")
80
+ if node.fallback and node.fallback not in self._nodes_map:
81
+ raise ValueError(f"Fallback node '{node.fallback}' for stage {node.id} not found in pipeline")
82
+ case ConditionNode():
83
+ for cond in node.conditions:
84
+ if cond.then_goto not in self._nodes_map:
85
+ raise ValueError(f"Condition target '{cond.then_goto}' not found in pipeline")
86
+ if node.else_goto and node.else_goto not in self._nodes_map:
87
+ raise ValueError(f"Condition else '{node.else_goto}' not found in pipeline")
88
+ case ParallelNode():
89
+ for child in node.children:
90
+ if child not in self._nodes_map:
91
+ raise ValueError(f"Parallel branch '{child}' not found in pipeline")
92
+ if node.next and node.next not in self._nodes_map:
93
+ raise ValueError(f"Parallel next '{node.next}' not found in pipeline")
94
+ case TerminalNode():
95
+ pass
96
+ case SubPipelineNode():
97
+ if node.subpipeline_id not in self.subpipelines:
98
+ raise ValueError(f"Subpipeline '{node.subpipeline_id}' not found for node {node.id}")
99
+ # Basic cycle check: prevent direct self-reference
100
+ if node.subpipeline_id == self.entry:
101
+ raise ValueError(f"Subpipeline '{node.subpipeline_id}' cannot reference root entry")
102
+ case MapNode():
103
+ if node.next and node.next not in self._nodes_map:
104
+ raise ValueError(f"Map next '{node.next}' not found in pipeline")
105
+ case _:
106
+ raise ValueError(f"Unknown node type: {type(node)}")
107
+
108
+ return True