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.
- stageflow/__init__.py +25 -0
- stageflow/builtins/__init__.py +34 -0
- stageflow/builtins/dicts.py +59 -0
- stageflow/builtins/lists.py +71 -0
- stageflow/builtins/lists_extra.py +110 -0
- stageflow/builtins/logic.py +84 -0
- stageflow/builtins/strings.py +81 -0
- stageflow/builtins/vars.py +118 -0
- stageflow/core/__init__.py +7 -0
- stageflow/core/context.py +86 -0
- stageflow/core/event.py +71 -0
- stageflow/core/jsonlogic.py +43 -0
- stageflow/core/node.py +256 -0
- stageflow/core/pipeline.py +108 -0
- stageflow/core/session.py +398 -0
- stageflow/core/stage.py +152 -0
- stageflow/core/utils.py +62 -0
- stageflow/docs/__init__.py +4 -0
- stageflow/docs/html.py +479 -0
- stageflow/docs/schema.py +38 -0
- stageflow/docs/schemas/pipeline.json +134 -0
- stageflow/py.typed +1 -0
- stageflow/testing.py +59 -0
- stageflow_framework-0.1.0.dist-info/METADATA +167 -0
- stageflow_framework-0.1.0.dist-info/RECORD +27 -0
- stageflow_framework-0.1.0.dist-info/WHEEL +5 -0
- stageflow_framework-0.1.0.dist-info/top_level.txt +1 -0
|
@@ -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
|
stageflow/core/stage.py
ADDED
|
@@ -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
|
+
}
|
stageflow/core/utils.py
ADDED
|
@@ -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__}")
|