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,398 @@
1
+ import asyncio
2
+ import copy
3
+ from typing import Callable, Any
4
+
5
+ from .context import Context
6
+ from .event import Event
7
+ from .pipeline import Pipeline
8
+ from .node import StageNode, ConditionNode, ParallelNode, TerminalNode, SubPipelineNode
9
+
10
+
11
+ class SessionResult:
12
+ artifacts: dict[str, dict | None]
13
+ result: dict | None
14
+ history: list[Event]
15
+ context: Context
16
+
17
+ def __init__(self, artifacts: dict, result: dict | None, history: list[Event], context: Context):
18
+ self.artifacts = artifacts
19
+ self.result = result
20
+ self.history = history
21
+ self.context = context
22
+
23
+ def to_dict(self) -> dict:
24
+ return {
25
+ "artifacts": self.artifacts,
26
+ "result": self.result,
27
+ "history": [e.to_dict() for e in self.history],
28
+ "context": self.context.to_dict(),
29
+ }
30
+
31
+
32
+ class Session:
33
+ def __init__(
34
+ self,
35
+ id: str,
36
+ pipeline: Pipeline,
37
+ context: Context = None,
38
+ event_handler: Callable[[Event], None] = None,
39
+ ):
40
+ pipeline.validate()
41
+ self.id = id
42
+ self.pipeline = pipeline
43
+ self.context = context or Context()
44
+ self.event_handler = event_handler or (lambda event: None)
45
+
46
+ self.artifact_paths: list[str] = []
47
+ self.result: dict | None = None
48
+ self.event_history: list[Event] = []
49
+
50
+ self.input_history: list[dict[str, Any]] = []
51
+ self._waiting: dict[str, list[asyncio.Future]] = {}
52
+ self._pending_inputs: dict[str, list[dict[str, Any]]] = {}
53
+
54
+ self._stopped = False
55
+ self._paused = False
56
+ self._skip_requested = False
57
+ self._current_node_id: str | None = None
58
+
59
+ def emit(self, event: Event):
60
+ self.event_history.append(event)
61
+ self.event_handler(event)
62
+
63
+ async def input(self, type_: str, payload: dict[str, Any]):
64
+ entry = {"type": type_, "payload": payload}
65
+ self.input_history.append(entry)
66
+ self.emit(Event(type="user_input", session_id=self.id, payload=entry))
67
+
68
+ if type_ == "command":
69
+ cmd = payload.get("name")
70
+ if cmd == "stop":
71
+ self._stopped = True
72
+ self.emit(Event(type="session_stopped", session_id=self.id))
73
+ elif cmd == "skip":
74
+ self._skip_requested = True
75
+ elif cmd == "pause":
76
+ self._paused = True
77
+ self.emit(Event(type="session_paused", session_id=self.id))
78
+ elif cmd == "resume":
79
+ self._paused = False
80
+ self.emit(Event(type="session_resumed", session_id=self.id))
81
+
82
+ if type_ in self._waiting:
83
+ for fut in list(self._waiting[type_]):
84
+ if not fut.done():
85
+ fut.set_result(entry)
86
+ else:
87
+ self._pending_inputs.setdefault(type_, []).append(entry)
88
+ return entry
89
+
90
+ async def wait_input(self, type_: str, timeout: float | None = None):
91
+ # Deliver buffered input if it arrived before waiter.
92
+ pending = self._pending_inputs.get(type_)
93
+ if pending:
94
+ return pending.pop(0)
95
+ loop = asyncio.get_running_loop()
96
+ fut = loop.create_future()
97
+ self._waiting.setdefault(type_, []).append(fut)
98
+ self.emit(Event(type="waiting_for_input", session_id=self.id, payload={"type": type_}))
99
+ try:
100
+ return await asyncio.wait_for(fut, timeout=timeout)
101
+ except asyncio.TimeoutError:
102
+ self.emit(Event(type="input_timeout", session_id=self.id, payload={"type": type_}))
103
+ return None
104
+ finally:
105
+ waiters = self._waiting.get(type_)
106
+ if waiters and fut in waiters:
107
+ waiters.remove(fut)
108
+ if not waiters:
109
+ del self._waiting[type_]
110
+
111
+ def last_input(self, type_: str | None = None):
112
+ if not self.input_history:
113
+ return None
114
+ if type_ is None:
115
+ return self.input_history[-1]
116
+ for entry in reversed(self.input_history):
117
+ if entry["type"] == type_:
118
+ return entry
119
+ return None
120
+
121
+ def _merge_artifacts(self, buffer: dict, new_values: dict):
122
+ for path, value in new_values.items():
123
+ if path in buffer and buffer[path] != value:
124
+ raise RuntimeError(f"Artifact merge conflict at '{path}'")
125
+ buffer[path] = value
126
+
127
+ def snapshot(self) -> dict:
128
+ return {
129
+ "session_id": self.id,
130
+ "current_node_id": self._current_node_id,
131
+ "result": self.result,
132
+ "artifact_paths": list(self.artifact_paths),
133
+ "context": self.context.to_dict(),
134
+ "event_history": [e.to_dict() for e in self.event_history],
135
+ "input_history": list(self.input_history),
136
+ "stopped": self._stopped,
137
+ "paused": self._paused,
138
+ "skip_requested": self._skip_requested,
139
+ "pipeline": self.pipeline.raw_json,
140
+ }
141
+
142
+ @classmethod
143
+ def from_snapshot(
144
+ cls,
145
+ snapshot: dict,
146
+ pipeline: Pipeline | None = None,
147
+ event_handler: Callable[[Event], None] | None = None,
148
+ ) -> "Session":
149
+ pipe = pipeline
150
+ if pipe is None:
151
+ raw = snapshot.get("pipeline")
152
+ if raw:
153
+ pipe = Pipeline.from_dict(raw)
154
+ if pipe is None:
155
+ raise ValueError("Pipeline is required to restore session")
156
+ ctx = Context.from_dict(snapshot.get("context", {}))
157
+ sess = cls(
158
+ id=snapshot.get("session_id"),
159
+ pipeline=pipe,
160
+ context=ctx,
161
+ event_handler=event_handler,
162
+ )
163
+ sess._current_node_id = snapshot.get("current_node_id")
164
+ sess.result = snapshot.get("result")
165
+ sess.artifact_paths = snapshot.get("artifact_paths", [])
166
+ sess.input_history = snapshot.get("input_history", [])
167
+ sess._stopped = snapshot.get("stopped", False)
168
+ sess._paused = snapshot.get("paused", False)
169
+ sess._skip_requested = snapshot.get("skip_requested", False)
170
+ sess.event_history = [Event.from_dict(e) for e in snapshot.get("event_history", [])]
171
+ return sess
172
+
173
+ async def run(self) -> SessionResult:
174
+ self.emit(Event(type="session_started", session_id=self.id))
175
+ node = self.pipeline.get_node(self._current_node_id) if self._current_node_id else self.pipeline.get_entry_node()
176
+
177
+ while node:
178
+ self._current_node_id = node.id
179
+ while self._paused and not self._stopped:
180
+ await asyncio.sleep(0.1)
181
+
182
+ if self._stopped:
183
+ self.result = {"result": "stopped"}
184
+ break
185
+
186
+ if self._skip_requested and isinstance(node, StageNode):
187
+ stage_class = node.get_stage_class()
188
+ if getattr(stage_class, "skipable", True):
189
+ self.emit(Event(type="stage_skipped", session_id=self.id, stage_id=node.id))
190
+ node = self.pipeline.get_node(node.next) if node.next else None
191
+ self._skip_requested = False
192
+ continue
193
+ else:
194
+ self.emit(Event(type="skip_denied", session_id=self.id, stage_id=node.id))
195
+ self._skip_requested = False
196
+
197
+ handler = self._get_handler(node)
198
+ handler_task = asyncio.create_task(handler(node))
199
+ try:
200
+ while not handler_task.done():
201
+ if self._stopped:
202
+ handler_task.cancel()
203
+ self.result = {"result": "stopped"}
204
+ try:
205
+ await handler_task
206
+ except asyncio.CancelledError:
207
+ pass
208
+ break
209
+ await asyncio.sleep(0.05)
210
+ if handler_task.cancelled():
211
+ break
212
+ next_node = await handler_task
213
+ except asyncio.CancelledError:
214
+ self.result = {"result": "stopped"}
215
+ break
216
+ except Exception as e:
217
+ raise e
218
+
219
+ if next_node is None and not isinstance(node, TerminalNode):
220
+ raise RuntimeError(f"No next node for {node.id} in session {self.id}")
221
+
222
+ node = next_node
223
+
224
+ self._current_node_id = None
225
+ self.emit(Event(type="session_completed", session_id=self.id))
226
+ artifacts = {
227
+ a: self.context.get(a, None)
228
+ for a in self.artifact_paths
229
+ }
230
+ return SessionResult(
231
+ artifacts=artifacts,
232
+ result=self.result,
233
+ history=self.event_history,
234
+ context=self.context,
235
+ )
236
+
237
+ def _get_handler(self, node):
238
+ if isinstance(node, TerminalNode):
239
+ return self._handle_terminal
240
+ if isinstance(node, StageNode):
241
+ return self._handle_stage
242
+ if isinstance(node, ConditionNode):
243
+ return self._handle_condition
244
+ if isinstance(node, ParallelNode):
245
+ return self._handle_parallel
246
+ if isinstance(node, SubPipelineNode):
247
+ return self._handle_subpipeline
248
+ raise ValueError(f"Unknown node type: {type(node)}")
249
+
250
+ async def _handle_terminal(self, node: TerminalNode):
251
+ self.artifact_paths = node.artifact_paths
252
+ self.result = node.result
253
+ self.emit(Event(type="session_terminated", session_id=self.id, stage_id=node.id))
254
+ return None
255
+
256
+ async def _handle_stage(self, node: StageNode):
257
+ stage_class = node.get_stage_class()
258
+ if stage_class is None:
259
+ raise ValueError(f"Stage class for node {node.id} not found")
260
+ stage_instance = stage_class(stage_id=node.id, config=node.config,
261
+ arguments=node.arguments, outputs=node.outputs, session=self)
262
+ errors = []
263
+ for retry in range(stage_instance.retries + 1):
264
+ try:
265
+ await asyncio.wait_for(stage_instance.run(), timeout=stage_instance.timeout)
266
+ except asyncio.TimeoutError as e:
267
+ self.emit(Event(type="stage_timeout", session_id=self.id, stage_id=node.id))
268
+ errors.append("timeout: " + str(e))
269
+ continue
270
+ except Exception as e:
271
+ self.emit(Event(type="stage_failed", session_id=self.id, stage_id=node.id, payload={"error": str(e)}))
272
+ errors.append(str(e))
273
+ continue
274
+ return self.pipeline.get_node(node.next) if node.next else None
275
+ if node.fallback:
276
+ return self.pipeline.get_node(node.fallback)
277
+ raise RuntimeError(f"Stage {node.id} failed after retries: {errors}")
278
+
279
+ async def _handle_condition(self, node: ConditionNode):
280
+ next_node = None
281
+ for condition in node.conditions:
282
+ if condition.if_condition.evaluate(self.context):
283
+ next_node = self.pipeline.get_node(condition.then_goto)
284
+ break
285
+ if not next_node and node.else_goto:
286
+ next_node = self.pipeline.get_node(node.else_goto)
287
+
288
+ self.emit(Event(
289
+ type="condition_evaluated",
290
+ session_id=self.id,
291
+ stage_id=node.id,
292
+ payload={"next_node": next_node.id if next_node else None},
293
+ ))
294
+ return next_node
295
+
296
+ async def _handle_parallel(self, node: ParallelNode):
297
+ branch_tasks = {
298
+ branch_id: asyncio.create_task(self._run_branch_graph(branch_id))
299
+ for branch_id in node.children
300
+ }
301
+ errors = []
302
+
303
+ async def cancel_pending(pending):
304
+ for t in pending:
305
+ t.cancel()
306
+ for t in pending:
307
+ try:
308
+ await t
309
+ except Exception:
310
+ pass
311
+
312
+ if node.policy == "all":
313
+ results = await asyncio.gather(*branch_tasks.values(), return_exceptions=True)
314
+ for branch_id, res in zip(branch_tasks.keys(), results):
315
+ if isinstance(res, Exception):
316
+ errors.append((branch_id, res))
317
+ if node.cancel_on_error:
318
+ break
319
+ if errors and node.cancel_on_error:
320
+ raise RuntimeError(f"Parallel node {node.id} failed: {errors}")
321
+ elif node.policy == "any":
322
+ pending = set(branch_tasks.values())
323
+ success = False
324
+ while pending and not success:
325
+ done, pending = await asyncio.wait(pending, return_when=asyncio.FIRST_COMPLETED)
326
+ for t in done:
327
+ try:
328
+ await t
329
+ success = True
330
+ break
331
+ except Exception as e:
332
+ errors.append(("unknown", e))
333
+ if node.cancel_on_error:
334
+ await cancel_pending(pending)
335
+ raise RuntimeError(f"Parallel node {node.id} failed: {errors}")
336
+ await cancel_pending(pending)
337
+ if not success:
338
+ raise RuntimeError(f"Parallel node {node.id} completed with no successful branch")
339
+
340
+ next_node = self.pipeline.get_node(node.next) if node.next else None
341
+ self.emit(Event(
342
+ type="parallel_completed",
343
+ session_id=self.id,
344
+ stage_id=node.id,
345
+ payload={"next_node": next_node.id if next_node else None},
346
+ ))
347
+ return next_node
348
+
349
+ async def _run_branch_graph(self, start_id: str) -> None:
350
+ node = self.pipeline.get_node(start_id)
351
+ while node:
352
+ if isinstance(node, TerminalNode):
353
+ return
354
+ handler = self._get_handler(node)
355
+ if handler is self._handle_terminal:
356
+ return
357
+ next_node = await handler(node)
358
+ node = next_node
359
+
360
+ async def _run_subpipeline(self, node: SubPipelineNode):
361
+ if node.subpipeline_id not in self.pipeline.subpipelines:
362
+ raise ValueError(f"Subpipeline '{node.subpipeline_id}' not found")
363
+ subpipeline_data = dict(self.pipeline.subpipelines[node.subpipeline_id])
364
+ subpipeline_data["subpipelines"] = self.pipeline.subpipelines
365
+ subpipeline = Pipeline.from_dict(subpipeline_data)
366
+ child_ctx = Context(payload={})
367
+ for child_path, parent_path in node.inputs.items():
368
+ child_ctx.set(child_path, copy.deepcopy(self.context.get(parent_path)))
369
+
370
+ def proxy_event(event: Event):
371
+ event.payload = {"subpipeline_node": node.id, **(event.payload or {})}
372
+ self.emit(event)
373
+
374
+ child_session = Session(
375
+ id=f"{self.id}:{node.id}",
376
+ pipeline=subpipeline,
377
+ context=child_ctx,
378
+ event_handler=proxy_event,
379
+ )
380
+ result = await child_session.run()
381
+ for parent_path, child_art in node.artifact_outputs.items():
382
+ if child_art not in result.artifacts:
383
+ raise ValueError(f"Artifact '{child_art}' not found in subpipeline '{node.subpipeline_id}'")
384
+ self.context.set(parent_path, result.artifacts[child_art])
385
+ if node.result_output:
386
+ self.context.set(node.result_output, result.result)
387
+ return {"artifacts": result.artifacts, "result": result.result}
388
+
389
+ async def _handle_subpipeline(self, node: SubPipelineNode):
390
+ await self._run_subpipeline(node)
391
+ next_node = self.pipeline.get_node(node.next) if node.next else None
392
+ self.emit(Event(
393
+ type="subpipeline_completed",
394
+ session_id=self.id,
395
+ stage_id=node.id,
396
+ payload={"next_node": next_node.id if next_node else None, "subpipeline_id": node.subpipeline_id},
397
+ ))
398
+ return next_node
@@ -0,0 +1,152 @@
1
+ import copy
2
+ from typing import Any, TYPE_CHECKING
3
+ import yaml
4
+ from stageflow.core import EventSpec, InputSpec
5
+ from .utils import validate_schema
6
+
7
+
8
+ STAGE_REGISTRY: dict[str, type["BaseStage"]] = {}
9
+
10
+ if TYPE_CHECKING:
11
+ from .session import Session
12
+
13
+
14
+ def register_stage(name: str):
15
+ def decorator(cls: type["BaseStage"]):
16
+ if name in STAGE_REGISTRY:
17
+ raise ValueError(f"Stage '{name}' already registered")
18
+ cls.stage_name = name
19
+ STAGE_REGISTRY[name] = cls
20
+ return cls
21
+ return decorator
22
+
23
+
24
+ def get_stage(name: str) -> type["BaseStage"]:
25
+ if name not in STAGE_REGISTRY:
26
+ raise ValueError(f"Stage '{name}' not found in registry")
27
+ return STAGE_REGISTRY[name]
28
+
29
+
30
+ def get_stages() -> dict[str, Any]:
31
+ return STAGE_REGISTRY
32
+
33
+
34
+ def get_stages_by_category() -> dict[str, list[type["BaseStage"]]]:
35
+ categories: dict[str, list[type["BaseStage"]]] = {}
36
+ for stage in STAGE_REGISTRY.values():
37
+ category = stage.category or "default"
38
+ categories.setdefault(category, []).append(stage)
39
+ return categories
40
+
41
+
42
+ def _normalize_field_entry(name: str | None, spec: Any) -> dict[str, Any]:
43
+ if isinstance(spec, dict):
44
+ entry = {"name": name, **spec} if name else dict(spec)
45
+ else:
46
+ entry = {"name": name, "type": spec}
47
+ entry.setdefault("type", "any")
48
+ entry.setdefault("optional", False)
49
+ entry.setdefault("description", "")
50
+ return entry
51
+
52
+
53
+ def _normalize_fields(raw: Any) -> list[dict[str, Any]]:
54
+ if not raw:
55
+ return []
56
+ normalized: list[dict[str, Any]] = []
57
+ if isinstance(raw, dict):
58
+ for name, spec in raw.items():
59
+ normalized.append(_normalize_field_entry(name, spec))
60
+ return normalized
61
+
62
+ if isinstance(raw, list):
63
+ for item in raw:
64
+ if isinstance(item, dict):
65
+ if "name" in item:
66
+ normalized.append(_normalize_field_entry(item.get("name"), {k: v for k, v in item.items() if k != "name"}))
67
+ elif len(item) == 1:
68
+ name, spec = next(iter(item.items()))
69
+ normalized.append(_normalize_field_entry(name, spec))
70
+ else:
71
+ normalized.append(_normalize_field_entry(str(item), {"type": "any"}))
72
+ return normalized
73
+
74
+ return normalized
75
+
76
+
77
+ class BaseStage:
78
+ skipable: bool = False
79
+ stage_name: str = "BaseStage"
80
+ category: str | None = None
81
+ allowed_events: list[EventSpec] = []
82
+ allowed_inputs: list[InputSpec] = []
83
+ timeout: float | None = 30
84
+ retries: int = 0
85
+
86
+ def __init__(self, stage_id: str, config: dict, arguments: dict, outputs: dict, session: "Session"):
87
+ self.stage_id = stage_id
88
+ self.config = config or {}
89
+ self.arguments_paths = arguments or {}
90
+ self.outputs_paths = outputs or {}
91
+ self.session = session
92
+
93
+ def get_arguments(self) -> dict:
94
+ arguments = dict()
95
+ for key, path in self.arguments_paths.items():
96
+ arguments[key] = copy.deepcopy(self.session.context.get(path))
97
+ return arguments
98
+
99
+ def set_outputs(self, outputs: dict):
100
+ for key, value in outputs.items():
101
+ if key in self.outputs_paths:
102
+ path = self.outputs_paths[key]
103
+ self.session.context.set(path, value)
104
+
105
+ async def run(self):
106
+ raise NotImplementedError
107
+
108
+ def emit(self, event_type: str, payload: dict | None = None):
109
+ from .event import Event
110
+ if self.allowed_events:
111
+ allowed = {spec.type for spec in self.allowed_events if spec.type}
112
+ if allowed and event_type not in allowed:
113
+ raise ValueError(f"Event type '{event_type}' is not allowed for stage '{self.stage_name}'")
114
+ matching = next((spec for spec in self.allowed_events if spec.type == event_type), None)
115
+ if matching and matching.payload_schema is not None:
116
+ validate_schema(payload or {}, matching.payload_schema, "Event payload")
117
+ self.session.emit(Event(
118
+ type=event_type,
119
+ session_id=self.session.id,
120
+ stage_id=self.stage_id,
121
+ payload=payload or {},
122
+ ))
123
+
124
+ async def wait_input(self, type_: str, timeout: float | None = None):
125
+ if self.allowed_inputs:
126
+ allowed = {spec.type for spec in self.allowed_inputs if spec.type}
127
+ if allowed and type_ not in allowed:
128
+ raise ValueError(f"Input type '{type_}' is not allowed for stage '{self.stage_name}'")
129
+ matching = next((spec for spec in self.allowed_inputs if spec.type == type_), None)
130
+ else:
131
+ matching = None
132
+ result = await self.session.wait_input(type_, timeout=timeout)
133
+ if result is None:
134
+ return None
135
+ if matching and matching.payload_schema is not None:
136
+ validate_schema(result.get("payload", {}), matching.payload_schema, "Input payload")
137
+ return result
138
+
139
+ @classmethod
140
+ def get_specs(cls) -> dict[str, Any]:
141
+ parsed_description = yaml.safe_load(cls.__doc__) if cls.__doc__ else {}
142
+ return {
143
+ "stage_name": cls.stage_name,
144
+ "skipable": cls.skipable,
145
+ "allowed_events": [e.__dict__ for e in cls.allowed_events],
146
+ "allowed_inputs": [i.__dict__ for i in cls.allowed_inputs],
147
+ "category": cls.category,
148
+ "description": parsed_description.get("description", "") if isinstance(parsed_description, dict) else "",
149
+ "arguments": _normalize_fields(parsed_description.get("arguments", [])) if isinstance(parsed_description, dict) else [],
150
+ "config": _normalize_fields(parsed_description.get("config", [])) if isinstance(parsed_description, dict) else [],
151
+ "outputs": _normalize_fields(parsed_description.get("outputs", [])) if isinstance(parsed_description, dict) else [],
152
+ }
@@ -0,0 +1,62 @@
1
+ from typing import Any, get_origin, get_args, Union
2
+
3
+
4
+ def validate_schema(value: Any, schema: object, path: str = "payload"):
5
+ if schema is None:
6
+ return
7
+
8
+ origin = get_origin(schema)
9
+ args = get_args(schema)
10
+
11
+ if schema is Any or schema is object:
12
+ return
13
+
14
+ if origin is Union and args:
15
+ for variant in args:
16
+ try:
17
+ validate_schema(value, variant, path)
18
+ return
19
+ except ValueError:
20
+ continue
21
+ raise ValueError(f"{path} expected one of {args}, got {type(value).__name__}")
22
+
23
+ if origin is list and args:
24
+ if not isinstance(value, list):
25
+ raise ValueError(f"{path} expected list, got {type(value).__name__}")
26
+ elem_schema = args[0]
27
+ for idx, item in enumerate(value):
28
+ validate_schema(item, elem_schema, f"{path}[{idx}]")
29
+ return
30
+ if origin is dict and args:
31
+ key_schema, val_schema = args
32
+ if not isinstance(value, dict):
33
+ raise ValueError(f"{path} expected dict, got {type(value).__name__}")
34
+ for k, v in value.items():
35
+ validate_schema(k, key_schema, f"{path} (key)")
36
+ validate_schema(v, val_schema, f"{path}[{k}]")
37
+ return
38
+
39
+ if isinstance(schema, dict):
40
+ if not isinstance(value, dict):
41
+ raise ValueError(f"{path} expected dict, got {type(value).__name__}")
42
+ for k, sub_schema in schema.items():
43
+ if k not in value:
44
+ raise ValueError(f"{path} missing required field '{k}'")
45
+ validate_schema(value[k], sub_schema, f"{path}.{k}")
46
+ return
47
+
48
+ if isinstance(schema, list) and len(schema) == 1:
49
+ if not isinstance(value, list):
50
+ raise ValueError(f"{path} expected list, got {type(value).__name__}")
51
+ for idx, item in enumerate(value):
52
+ validate_schema(item, schema[0], f"{path}[{idx}]")
53
+ return
54
+
55
+ if isinstance(schema, tuple):
56
+ expected = schema
57
+ elif isinstance(schema, type):
58
+ expected = (schema,)
59
+ else:
60
+ return
61
+ if not isinstance(value, expected):
62
+ raise ValueError(f"{path} expected {schema}, got {type(value).__name__}")
@@ -0,0 +1,4 @@
1
+ from .schema import generate_stages_yaml, generate_stages_json # noqa
2
+
3
+
4
+ __all__ = ["generate_stages_yaml", "generate_stages_json"]