harness-sdk-python 0.4.0__tar.gz → 0.4.2__tar.gz
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.
- {harness_sdk_python-0.4.0 → harness_sdk_python-0.4.2}/PKG-INFO +1 -1
- {harness_sdk_python-0.4.0 → harness_sdk_python-0.4.2}/pyproject.toml +1 -1
- {harness_sdk_python-0.4.0 → harness_sdk_python-0.4.2}/src/harness_sdk/run_manager.py +126 -70
- {harness_sdk_python-0.4.0 → harness_sdk_python-0.4.2}/tests/run_helpers.py +16 -2
- {harness_sdk_python-0.4.0 → harness_sdk_python-0.4.2}/tests/test_batches.py +2 -1
- {harness_sdk_python-0.4.0 → harness_sdk_python-0.4.2}/tests/test_branch_anchor.py +8 -16
- {harness_sdk_python-0.4.0 → harness_sdk_python-0.4.2}/tests/test_caller.py +10 -1
- {harness_sdk_python-0.4.0 → harness_sdk_python-0.4.2}/tests/test_edit_dispatched.py +5 -0
- {harness_sdk_python-0.4.0 → harness_sdk_python-0.4.2}/tests/test_edit_reload.py +13 -9
- {harness_sdk_python-0.4.0 → harness_sdk_python-0.4.2}/tests/test_enqueue.py +13 -1
- {harness_sdk_python-0.4.0 → harness_sdk_python-0.4.2}/tests/test_facade.py +6 -3
- {harness_sdk_python-0.4.0 → harness_sdk_python-0.4.2}/tests/test_input_required.py +11 -2
- {harness_sdk_python-0.4.0 → harness_sdk_python-0.4.2}/tests/test_outcomes.py +13 -2
- {harness_sdk_python-0.4.0 → harness_sdk_python-0.4.2}/tests/test_rewind_during_run.py +9 -4
- harness_sdk_python-0.4.2/tests/test_run_leaf.py +121 -0
- harness_sdk_python-0.4.2/tests/test_settle.py +218 -0
- {harness_sdk_python-0.4.0 → harness_sdk_python-0.4.2}/tests/test_steer.py +5 -0
- {harness_sdk_python-0.4.0 → harness_sdk_python-0.4.2}/tests/test_stop_continue.py +32 -5
- harness_sdk_python-0.4.0/tests/test_settle.py +0 -131
- {harness_sdk_python-0.4.0 → harness_sdk_python-0.4.2}/.gitignore +0 -0
- {harness_sdk_python-0.4.0 → harness_sdk_python-0.4.2}/README.md +0 -0
- {harness_sdk_python-0.4.0 → harness_sdk_python-0.4.2}/src/harness_sdk/__init__.py +0 -0
- {harness_sdk_python-0.4.0 → harness_sdk_python-0.4.2}/src/harness_sdk/fenced_postgres.py +0 -0
- {harness_sdk_python-0.4.0 → harness_sdk_python-0.4.2}/tests/test_fenced_postgres.py +0 -0
- {harness_sdk_python-0.4.0 → harness_sdk_python-0.4.2}/tests/test_placement.py +0 -0
|
@@ -6,8 +6,14 @@ Implements the full state x command matrix from runs.mdx on a single
|
|
|
6
6
|
Commands stage entries into an intake list and schedule a drain; the drain
|
|
7
7
|
reduces the whole intake into the lanes in order (pure placement), then takes
|
|
8
8
|
exactly one action based on status and net effect (interrupt, dispatch, or
|
|
9
|
-
continue).
|
|
9
|
+
continue). A parked send settles at staging; run-starting commands settle at
|
|
10
|
+
the executor's ``ctx.ack_messages()``; ``run/stop`` awaits a future the drain
|
|
10
11
|
resolves once the in-flight run has ended.
|
|
12
|
+
|
|
13
|
+
Dispatched entries leave the queue projection but are retained until the
|
|
14
|
+
executor acks. A run that ends before the ack puts them back at the front of
|
|
15
|
+
their lane and settles their sends accepted; edit, reload, and continue
|
|
16
|
+
initiators settle rejected (``stopped`` or the failure).
|
|
11
17
|
"""
|
|
12
18
|
|
|
13
19
|
import asyncio
|
|
@@ -90,7 +96,6 @@ class _Edit:
|
|
|
90
96
|
@dataclass
|
|
91
97
|
class _Reload:
|
|
92
98
|
source_meta: dict[str, Any]
|
|
93
|
-
message_id: str
|
|
94
99
|
ack: Callable[[], None]
|
|
95
100
|
future: "asyncio.Future[Any]" = field(default_factory=_future)
|
|
96
101
|
|
|
@@ -102,7 +107,7 @@ class _Stop:
|
|
|
102
107
|
|
|
103
108
|
@dataclass
|
|
104
109
|
class _Continue:
|
|
105
|
-
|
|
110
|
+
future: "asyncio.Future[Any]" = field(default_factory=_future)
|
|
106
111
|
|
|
107
112
|
|
|
108
113
|
@dataclass
|
|
@@ -119,7 +124,6 @@ class _Rewind:
|
|
|
119
124
|
rollback_to: Any
|
|
120
125
|
ack: Callable[[], None]
|
|
121
126
|
future: "asyncio.Future[Any]"
|
|
122
|
-
message_id: str | None = None
|
|
123
127
|
acked: bool = False
|
|
124
128
|
|
|
125
129
|
|
|
@@ -128,6 +132,8 @@ class _Effects:
|
|
|
128
132
|
steer_added: bool = False
|
|
129
133
|
new_added: bool = False
|
|
130
134
|
continue_requested: bool = False
|
|
135
|
+
staged_sends: list[tuple[str, "asyncio.Future[Any]"]] = field(default_factory=list)
|
|
136
|
+
continues: list["asyncio.Future[Any]"] = field(default_factory=list)
|
|
131
137
|
|
|
132
138
|
|
|
133
139
|
class RunManager:
|
|
@@ -139,6 +145,7 @@ class RunManager:
|
|
|
139
145
|
get_message_meta: GetMessageMeta,
|
|
140
146
|
create_task: Callable[[Any], "asyncio.Task[Any]"],
|
|
141
147
|
schedule: Callable[[Callable[[], None]], None],
|
|
148
|
+
leaf_message_id: str | None,
|
|
142
149
|
capabilities: Iterable[str] = (),
|
|
143
150
|
max_queued: int = 50,
|
|
144
151
|
) -> None:
|
|
@@ -150,6 +157,10 @@ class RunManager:
|
|
|
150
157
|
raise ValueError("rewind-during-run requires the rewind capability")
|
|
151
158
|
if max_queued < 1:
|
|
152
159
|
raise ValueError("max_queued must be >= 1")
|
|
160
|
+
if leaf_message_id is not None and (
|
|
161
|
+
not isinstance(leaf_message_id, str) or leaf_message_id == ""
|
|
162
|
+
):
|
|
163
|
+
raise ValueError("leaf_message_id must be a non-empty string or None")
|
|
153
164
|
self._state = state
|
|
154
165
|
self._start = start
|
|
155
166
|
self._get_message_meta = get_message_meta
|
|
@@ -160,7 +171,6 @@ class RunManager:
|
|
|
160
171
|
self._task: "asyncio.Task[Any] | None" = None
|
|
161
172
|
self._ctx: "RunManager.StartContext | None" = None
|
|
162
173
|
self._dispatched_ids: tuple[str, ...] = ()
|
|
163
|
-
self._reload_message_id: str | None = None
|
|
164
174
|
self._dispatch_record: dict[str, Any] | None = None
|
|
165
175
|
self._callers: dict[str, StatewireClientHandle] = {}
|
|
166
176
|
self._intake: list[Any] = []
|
|
@@ -168,18 +178,22 @@ class RunManager:
|
|
|
168
178
|
self._staged_stops: list["asyncio.Future[Any]"] = []
|
|
169
179
|
self._stop_reason: str | None = None
|
|
170
180
|
self._staged_rewinds: list[_Rewind] = []
|
|
171
|
-
self.
|
|
181
|
+
self._dispatching: list[tuple[str, dict[str, Any]]] = []
|
|
182
|
+
self._run_acked = True
|
|
183
|
+
self._leaf_confirmed = leaf_message_id
|
|
184
|
+
self._send_futures: list["asyncio.Future[Any]"] = []
|
|
172
185
|
self._run_futures: list["asyncio.Future[Any]"] = []
|
|
173
186
|
self._input_requests: list[dict[str, Any]] = []
|
|
174
187
|
self._input_answers: dict[str, Any] = {}
|
|
175
|
-
self._init_state()
|
|
188
|
+
self._init_state(leaf_message_id)
|
|
176
189
|
|
|
177
|
-
def _init_state(self) -> None:
|
|
190
|
+
def _init_state(self, leaf_message_id: str | None) -> None:
|
|
178
191
|
# Runs state is write-only and not durable: overwrite whatever is there.
|
|
179
192
|
self._state["status"] = "ready"
|
|
180
193
|
self._state["error"] = None
|
|
181
194
|
self._state["queue"] = []
|
|
182
195
|
self._state["steerQueue"] = []
|
|
196
|
+
self._state["runLeafMessageId"] = leaf_message_id
|
|
183
197
|
self._state.pop("dispatch", None)
|
|
184
198
|
self._state.pop("inputRequests", None)
|
|
185
199
|
|
|
@@ -200,6 +214,9 @@ class RunManager:
|
|
|
200
214
|
def _continue_type(self) -> str:
|
|
201
215
|
return "error-continue" if self._status() == "error" else "stop-continue"
|
|
202
216
|
|
|
217
|
+
def _is_dispatching(self, message_id: str) -> bool:
|
|
218
|
+
return any(item["id"] == message_id for _, item in self._dispatching)
|
|
219
|
+
|
|
203
220
|
# ─── Staging and drain ──────────────────────────────────
|
|
204
221
|
|
|
205
222
|
async def _stage(self, entry: Any) -> Any:
|
|
@@ -221,15 +238,15 @@ class RunManager:
|
|
|
221
238
|
self._staged_stops.clear()
|
|
222
239
|
if isinstance(outcome, RunManager.Complete):
|
|
223
240
|
self._dispatched_ids = ()
|
|
224
|
-
self._reload_message_id = None
|
|
225
241
|
if self._staged_rewinds:
|
|
226
242
|
rewind = self._staged_rewinds.pop(0)
|
|
227
243
|
self._dispatch(
|
|
228
244
|
rewind.type,
|
|
229
245
|
rewind.messages,
|
|
230
246
|
rollback_to=rewind.rollback_to,
|
|
231
|
-
message_id=rewind.message_id,
|
|
232
247
|
)
|
|
248
|
+
for message in rewind.messages:
|
|
249
|
+
self._callers.pop(message["id"], None)
|
|
233
250
|
rewind.ack()
|
|
234
251
|
if not rewind.future.done():
|
|
235
252
|
self._run_futures.append(rewind.future)
|
|
@@ -244,6 +261,15 @@ class RunManager:
|
|
|
244
261
|
if not rewind.acked:
|
|
245
262
|
rewind.ack()
|
|
246
263
|
rewind.acked = True
|
|
264
|
+
dispatching = {item["id"] for _, item in self._dispatching}
|
|
265
|
+
for message_id, future in fx.staged_sends:
|
|
266
|
+
if message_id in dispatching:
|
|
267
|
+
self._send_futures.append(future)
|
|
268
|
+
elif not future.done():
|
|
269
|
+
future.set_result(None)
|
|
270
|
+
for future in fx.continues:
|
|
271
|
+
if not future.done():
|
|
272
|
+
future.set_result(None)
|
|
247
273
|
|
|
248
274
|
def _apply(self, entry: Any, fx: _Effects) -> None:
|
|
249
275
|
if isinstance(entry, _Stop):
|
|
@@ -251,6 +277,7 @@ class RunManager:
|
|
|
251
277
|
return
|
|
252
278
|
if isinstance(entry, _Continue):
|
|
253
279
|
fx.continue_requested = True
|
|
280
|
+
fx.continues.append(entry.future)
|
|
254
281
|
return
|
|
255
282
|
try:
|
|
256
283
|
if isinstance(entry, _Send):
|
|
@@ -277,6 +304,8 @@ class RunManager:
|
|
|
277
304
|
elif status in ("error", "stopped"):
|
|
278
305
|
if fx.continue_requested or fx.steer_added:
|
|
279
306
|
self._dispatch(self._continue_type(), [])
|
|
307
|
+
self._run_futures.extend(fx.continues)
|
|
308
|
+
fx.continues.clear()
|
|
280
309
|
elif fx.new_added and pre_empty:
|
|
281
310
|
self._pop_dispatchable()
|
|
282
311
|
elif status == "input-required" and self._input_requests and len(
|
|
@@ -321,7 +350,6 @@ class RunManager:
|
|
|
321
350
|
*,
|
|
322
351
|
rollback_to: Any = _ABSENT,
|
|
323
352
|
responses: Any = _ABSENT,
|
|
324
|
-
message_id: str | None = None,
|
|
325
353
|
) -> None:
|
|
326
354
|
if type not in _ENTRY_TYPES:
|
|
327
355
|
raise ValueError(f"invalid entry type: {type!r}")
|
|
@@ -331,22 +359,22 @@ class RunManager:
|
|
|
331
359
|
messages = [
|
|
332
360
|
{k: v for k, v in message.items() if k != "caller"} for message in messages
|
|
333
361
|
]
|
|
334
|
-
for message in messages:
|
|
335
|
-
self._callers.pop(message["id"], None)
|
|
336
|
-
self._adopt_entities(message["id"] for message in messages)
|
|
337
362
|
record: dict[str, Any] = {"type": type, "messages": list(messages)}
|
|
338
363
|
if rollback_to is not _ABSENT:
|
|
339
364
|
record["rollbackTo"] = rollback_to
|
|
340
365
|
if responses is not _ABSENT:
|
|
341
366
|
record["responses"] = responses
|
|
342
|
-
if message_id is not None:
|
|
343
|
-
record["messageId"] = message_id
|
|
344
367
|
self._dispatch_record = record
|
|
345
368
|
self._stop_reason = None
|
|
369
|
+
self._run_acked = False
|
|
370
|
+
self._leaf_confirmed = plain(self._state["runLeafMessageId"])
|
|
346
371
|
self._state["error"] = None
|
|
347
372
|
if messages:
|
|
348
373
|
self._dispatched_ids = tuple(m["id"] for m in messages)
|
|
349
|
-
|
|
374
|
+
if rollback_to is not _ABSENT:
|
|
375
|
+
self._state["runLeafMessageId"] = rollback_to
|
|
376
|
+
if messages:
|
|
377
|
+
self._state["runLeafMessageId"] = messages[-1]["id"]
|
|
350
378
|
self._state["status"] = "running"
|
|
351
379
|
ctx = RunManager.StartContext(
|
|
352
380
|
type=type,
|
|
@@ -356,7 +384,6 @@ class RunManager:
|
|
|
356
384
|
_manager=self,
|
|
357
385
|
_rollback_to=rollback_to,
|
|
358
386
|
_responses=responses,
|
|
359
|
-
_message_id=message_id,
|
|
360
387
|
)
|
|
361
388
|
self._ctx = ctx
|
|
362
389
|
self._task = self._create_task(self._run(ctx))
|
|
@@ -377,23 +404,34 @@ class RunManager:
|
|
|
377
404
|
"start must return a RunManager outcome, got "
|
|
378
405
|
f"{type(outcome).__name__}"
|
|
379
406
|
)
|
|
407
|
+
if isinstance(
|
|
408
|
+
outcome, (RunManager.Complete, RunManager.InputRequired)
|
|
409
|
+
) and (self._dispatching or not self._run_acked):
|
|
410
|
+
raise RuntimeError(
|
|
411
|
+
"run settled without acking its messages"
|
|
412
|
+
" (call ctx.ack_messages())"
|
|
413
|
+
)
|
|
380
414
|
except Exception as exc:
|
|
381
415
|
self._settle(ctx)
|
|
382
416
|
message = str(exc) or type(exc).__name__
|
|
383
417
|
if isinstance(exc, StatewireReject):
|
|
384
418
|
self._freeze(exc.message, exc.payload)
|
|
385
|
-
self.
|
|
419
|
+
self._settle_initiators(exc)
|
|
386
420
|
else:
|
|
387
421
|
self._freeze(message)
|
|
388
|
-
self.
|
|
422
|
+
self._settle_initiators(_reject("run-error", message))
|
|
423
|
+
self._revert_dispatching()
|
|
389
424
|
self._drain()
|
|
390
425
|
return
|
|
391
426
|
self._settle(ctx)
|
|
392
|
-
|
|
393
|
-
_reject("run-error", "run ended in error")
|
|
394
|
-
|
|
395
|
-
|
|
396
|
-
|
|
427
|
+
if isinstance(outcome, RunManager.Error):
|
|
428
|
+
error = _reject("run-error", "run ended in error")
|
|
429
|
+
elif isinstance(outcome, RunManager.Stop):
|
|
430
|
+
error = _reject("stopped", "run stopped before the messages-ack")
|
|
431
|
+
else:
|
|
432
|
+
error = None
|
|
433
|
+
self._settle_initiators(error)
|
|
434
|
+
self._revert_dispatching()
|
|
397
435
|
self._outcome = outcome
|
|
398
436
|
self._drain()
|
|
399
437
|
|
|
@@ -403,7 +441,7 @@ class RunManager:
|
|
|
403
441
|
self._task = None
|
|
404
442
|
self._dispatch_record = None
|
|
405
443
|
|
|
406
|
-
def
|
|
444
|
+
def _settle_initiators(self, error: StatewireReject | None) -> None:
|
|
407
445
|
futures, self._run_futures = self._run_futures, []
|
|
408
446
|
for future in futures:
|
|
409
447
|
if future.done():
|
|
@@ -413,21 +451,48 @@ class RunManager:
|
|
|
413
451
|
else:
|
|
414
452
|
future.set_exception(error)
|
|
415
453
|
|
|
416
|
-
def
|
|
417
|
-
|
|
418
|
-
|
|
419
|
-
if
|
|
420
|
-
|
|
454
|
+
def _settle_sends(self) -> None:
|
|
455
|
+
futures, self._send_futures = self._send_futures, []
|
|
456
|
+
for future in futures:
|
|
457
|
+
if not future.done():
|
|
458
|
+
future.set_result(None)
|
|
459
|
+
|
|
460
|
+
def _ack_messages(self) -> None:
|
|
461
|
+
taken, self._dispatching = self._dispatching, []
|
|
462
|
+
self._run_acked = True
|
|
463
|
+
for _, item in taken:
|
|
464
|
+
self._callers.pop(item["id"], None)
|
|
465
|
+
self._leaf_confirmed = plain(self._state["runLeafMessageId"])
|
|
466
|
+
self._settle_sends()
|
|
467
|
+
self._settle_initiators(None)
|
|
468
|
+
|
|
469
|
+
def _revert_dispatching(self) -> None:
|
|
470
|
+
self._settle_sends()
|
|
471
|
+
taken, self._dispatching = self._dispatching, []
|
|
472
|
+
if not taken:
|
|
473
|
+
return
|
|
474
|
+
for lane in ("steerQueue", "queue"):
|
|
475
|
+
front = [item for l, item in taken if l == lane]
|
|
476
|
+
if front:
|
|
477
|
+
self._state[lane] = front + self._lane_items(lane)
|
|
478
|
+
ids = {item["id"] for _, item in taken}
|
|
479
|
+
self._dispatched_ids = tuple(
|
|
480
|
+
id for id in self._dispatched_ids if id not in ids
|
|
481
|
+
)
|
|
482
|
+
if plain(self._state["runLeafMessageId"]) in ids:
|
|
483
|
+
self._state["runLeafMessageId"] = self._leaf_confirmed
|
|
421
484
|
|
|
422
485
|
def _pop_dispatchable(self) -> bool:
|
|
423
486
|
steer = self._lane_items("steerQueue")
|
|
424
487
|
if steer:
|
|
425
488
|
self._state["steerQueue"] = []
|
|
489
|
+
self._dispatching = [("steerQueue", item) for item in steer]
|
|
426
490
|
self._dispatch("message-send", steer)
|
|
427
491
|
return True
|
|
428
492
|
queue = self._lane_items("queue")
|
|
429
493
|
if queue:
|
|
430
494
|
self._state["queue"].pop(0)
|
|
495
|
+
self._dispatching = [("queue", queue[0])]
|
|
431
496
|
self._dispatch("message-send", [queue[0]])
|
|
432
497
|
return True
|
|
433
498
|
return False
|
|
@@ -529,7 +594,7 @@ class RunManager:
|
|
|
529
594
|
return
|
|
530
595
|
if self._lane_of(e.anchor) is not None:
|
|
531
596
|
return
|
|
532
|
-
if e.anchor in self._dispatched_ids
|
|
597
|
+
if e.anchor in self._dispatched_ids:
|
|
533
598
|
return
|
|
534
599
|
if e.anchor_meta is None:
|
|
535
600
|
raise _reject("unknown-id", f"anchor {e.anchor} names nothing")
|
|
@@ -619,9 +684,9 @@ class RunManager:
|
|
|
619
684
|
if e.lane == "steerQueue":
|
|
620
685
|
fx.steer_added = True
|
|
621
686
|
return None
|
|
622
|
-
if e.message_id in self._dispatched_ids:
|
|
687
|
+
if e.message_id in self._dispatched_ids or self._is_dispatching(e.message_id):
|
|
623
688
|
return self._park_dispatched_edit(e)
|
|
624
|
-
if e.meta is not None
|
|
689
|
+
if e.meta is not None:
|
|
625
690
|
raise _reject(
|
|
626
691
|
"duplicate-id", f"message id {e.message_id} is already used"
|
|
627
692
|
)
|
|
@@ -630,7 +695,7 @@ class RunManager:
|
|
|
630
695
|
self._check_anchor(e)
|
|
631
696
|
with self._caller_stamp(e.message_id, e.caller):
|
|
632
697
|
self._insert_new(e.lane, e.message, e.params)
|
|
633
|
-
|
|
698
|
+
fx.staged_sends.append((e.message_id, e.future))
|
|
634
699
|
e.ack()
|
|
635
700
|
fx.new_added = True
|
|
636
701
|
if e.lane == "steerQueue":
|
|
@@ -640,7 +705,9 @@ class RunManager:
|
|
|
640
705
|
def _apply_move(self, e: _Send, fx: _Effects) -> Any:
|
|
641
706
|
current = self._lane_of(e.message_id)
|
|
642
707
|
if current is None:
|
|
643
|
-
if e.message_id in self._dispatched_ids
|
|
708
|
+
if e.message_id in self._dispatched_ids or self._is_dispatching(
|
|
709
|
+
e.message_id
|
|
710
|
+
):
|
|
644
711
|
if e.lane == "steerQueue":
|
|
645
712
|
return None
|
|
646
713
|
raise _reject(
|
|
@@ -685,11 +752,6 @@ class RunManager:
|
|
|
685
752
|
self._state[lane] = [
|
|
686
753
|
item for item in self._lane_items(lane) if item["id"] != e.message_id
|
|
687
754
|
]
|
|
688
|
-
entity = self._entity_futures.pop(e.message_id, None)
|
|
689
|
-
if entity is not None and not entity.done():
|
|
690
|
-
entity.set_exception(
|
|
691
|
-
_reject("removed", f"message {e.message_id} was removed from the queue")
|
|
692
|
-
)
|
|
693
755
|
return None
|
|
694
756
|
|
|
695
757
|
def _apply_input(self, e: _Input) -> Any:
|
|
@@ -735,14 +797,6 @@ class RunManager:
|
|
|
735
797
|
|
|
736
798
|
def _apply_reload(self, e: _Reload) -> Any:
|
|
737
799
|
self._check_leaf_lanes(e.source_meta, "run/reload")
|
|
738
|
-
if (
|
|
739
|
-
self._lane_of(e.message_id) is not None
|
|
740
|
-
or e.message_id in self._dispatched_ids
|
|
741
|
-
or e.message_id == self._reload_message_id
|
|
742
|
-
):
|
|
743
|
-
raise _reject(
|
|
744
|
-
"duplicate-id", f"message id {e.message_id} is already used"
|
|
745
|
-
)
|
|
746
800
|
self._staged_rewinds.append(
|
|
747
801
|
_Rewind(
|
|
748
802
|
"message-reload",
|
|
@@ -750,7 +804,6 @@ class RunManager:
|
|
|
750
804
|
e.source_meta["parentId"],
|
|
751
805
|
e.ack,
|
|
752
806
|
e.future,
|
|
753
|
-
message_id=e.message_id,
|
|
754
807
|
)
|
|
755
808
|
)
|
|
756
809
|
return _PARKED
|
|
@@ -896,11 +949,6 @@ class RunManager:
|
|
|
896
949
|
source_id = params.get("sourceId") if isinstance(params, dict) else None
|
|
897
950
|
if not isinstance(source_id, str):
|
|
898
951
|
raise _reject("invalid-message", "sourceId must be a string")
|
|
899
|
-
message_id = params.get("messageId")
|
|
900
|
-
if not isinstance(message_id, str) or message_id == "":
|
|
901
|
-
raise _reject("invalid-message", "messageId must be a non-empty string")
|
|
902
|
-
if await self._get_message_meta(message_id) is not None:
|
|
903
|
-
raise _reject("duplicate-id", f"message id {message_id} is already used")
|
|
904
952
|
meta = await self._get_message_meta(source_id)
|
|
905
953
|
if meta is None:
|
|
906
954
|
raise _reject("unknown-id", f"message {source_id} is unknown")
|
|
@@ -918,7 +966,7 @@ class RunManager:
|
|
|
918
966
|
"capability-missing",
|
|
919
967
|
"the assistant-continuation capability is not enabled",
|
|
920
968
|
)
|
|
921
|
-
return await self._stage(_Reload(meta,
|
|
969
|
+
return await self._stage(_Reload(meta, ack))
|
|
922
970
|
|
|
923
971
|
async def stop(self, params: Any = None, *, ack: Callable[[], None]) -> Any:
|
|
924
972
|
if params is not None and not isinstance(params, dict):
|
|
@@ -986,7 +1034,7 @@ class RunManager:
|
|
|
986
1034
|
response = self._validated_response(request["type"], params["response"])
|
|
987
1035
|
return await self._stage(_Input(request_id, response))
|
|
988
1036
|
|
|
989
|
-
def continue_run(self) ->
|
|
1037
|
+
async def continue_run(self, *, ack: Callable[[], None]) -> Any:
|
|
990
1038
|
status = self._status()
|
|
991
1039
|
if status not in ("error", "stopped"):
|
|
992
1040
|
raise _reject("wrong-state", f"run/continue is rejected in {status}")
|
|
@@ -997,8 +1045,11 @@ class RunManager:
|
|
|
997
1045
|
"capability-missing",
|
|
998
1046
|
"bare continue requires the incomplete-continuation capability",
|
|
999
1047
|
)
|
|
1000
|
-
|
|
1048
|
+
entry = _Continue()
|
|
1049
|
+
self._intake.append(entry)
|
|
1050
|
+
ack()
|
|
1001
1051
|
self._schedule(self._drain)
|
|
1052
|
+
return await entry.future
|
|
1002
1053
|
|
|
1003
1054
|
# ─── Outcomes and context ───────────────────────────────
|
|
1004
1055
|
|
|
@@ -1052,20 +1103,11 @@ class RunManager:
|
|
|
1052
1103
|
_manager: "RunManager"
|
|
1053
1104
|
_rollback_to: Any
|
|
1054
1105
|
_responses: Any
|
|
1055
|
-
_message_id: str | None
|
|
1056
1106
|
|
|
1057
1107
|
@property
|
|
1058
1108
|
def stop_reason(self) -> str | None:
|
|
1059
1109
|
return self._manager._stop_reason
|
|
1060
1110
|
|
|
1061
|
-
@property
|
|
1062
|
-
def message_id(self) -> str:
|
|
1063
|
-
if self._message_id is None:
|
|
1064
|
-
raise AttributeError(
|
|
1065
|
-
"message_id is only present on message-reload entries"
|
|
1066
|
-
)
|
|
1067
|
-
return self._message_id
|
|
1068
|
-
|
|
1069
1111
|
@property
|
|
1070
1112
|
def has_rollback(self) -> bool:
|
|
1071
1113
|
return self._rollback_to is not _ABSENT
|
|
@@ -1098,13 +1140,27 @@ class RunManager:
|
|
|
1098
1140
|
self._ensure_active()
|
|
1099
1141
|
items = self._manager._lane_items("steerQueue")
|
|
1100
1142
|
self._manager._state["steerQueue"] = []
|
|
1101
|
-
for item in items
|
|
1102
|
-
|
|
1103
|
-
|
|
1143
|
+
self._manager._dispatching.extend(("steerQueue", item) for item in items)
|
|
1144
|
+
if items:
|
|
1145
|
+
self._manager._state["runLeafMessageId"] = items[-1]["id"]
|
|
1104
1146
|
return tuple(
|
|
1105
1147
|
{k: v for k, v in item.items() if k != "caller"} for item in items
|
|
1106
1148
|
)
|
|
1107
1149
|
|
|
1150
|
+
def ack_messages(self) -> None:
|
|
1151
|
+
self._ensure_active()
|
|
1152
|
+
self._manager._ack_messages()
|
|
1153
|
+
|
|
1154
|
+
def set_leaf_message_id(self, message_id: str) -> None:
|
|
1155
|
+
self._ensure_active()
|
|
1156
|
+
if not isinstance(message_id, str) or message_id == "":
|
|
1157
|
+
raise ValueError("message_id must be a non-empty string")
|
|
1158
|
+
if self._manager._dispatching or not self._manager._run_acked:
|
|
1159
|
+
raise RuntimeError(
|
|
1160
|
+
"ack_messages must precede set_leaf_message_id"
|
|
1161
|
+
)
|
|
1162
|
+
self._manager._state["runLeafMessageId"] = message_id
|
|
1163
|
+
|
|
1108
1164
|
def set_recovery_state(self, value: Any) -> None:
|
|
1109
1165
|
self._ensure_active()
|
|
1110
1166
|
record = self._manager._dispatch_record
|
|
@@ -15,6 +15,9 @@ class Call:
|
|
|
15
15
|
ctx: RunManager.StartContext
|
|
16
16
|
outcome: "asyncio.Future[Any]"
|
|
17
17
|
|
|
18
|
+
def ack(self) -> None:
|
|
19
|
+
self.ctx.ack_messages()
|
|
20
|
+
|
|
18
21
|
def finish(self, outcome: Any) -> None:
|
|
19
22
|
self.outcome.set_result(outcome)
|
|
20
23
|
|
|
@@ -40,6 +43,16 @@ class Script:
|
|
|
40
43
|
return {"isLeaf": not self.thread}
|
|
41
44
|
return self.thread.get(message_id)
|
|
42
45
|
|
|
46
|
+
def leaf_id(self) -> str | None:
|
|
47
|
+
return next(
|
|
48
|
+
(
|
|
49
|
+
id
|
|
50
|
+
for id, meta in self.thread.items()
|
|
51
|
+
if meta.get("isLeaf") and meta.get("onActiveBranch", True)
|
|
52
|
+
),
|
|
53
|
+
None,
|
|
54
|
+
)
|
|
55
|
+
|
|
43
56
|
async def next_call(self, timeout: float = 5) -> Call:
|
|
44
57
|
return await asyncio.wait_for(self.calls.get(), timeout)
|
|
45
58
|
|
|
@@ -62,6 +75,7 @@ def make_host(script: Script, capabilities=(), initial_runs=None, max_queued=50)
|
|
|
62
75
|
get_message_meta=script.get_message_meta,
|
|
63
76
|
create_task=self.create_task,
|
|
64
77
|
schedule=self.schedule,
|
|
78
|
+
leaf_message_id=script.leaf_id(),
|
|
65
79
|
capabilities=capabilities,
|
|
66
80
|
max_queued=max_queued,
|
|
67
81
|
)
|
|
@@ -92,8 +106,8 @@ def make_host(script: Script, capabilities=(), initial_runs=None, max_queued=50)
|
|
|
92
106
|
return await self.runs.stop(params, ack=ctx.ack)
|
|
93
107
|
|
|
94
108
|
@command("run/continue")
|
|
95
|
-
async def run_continue(self):
|
|
96
|
-
self.runs.continue_run()
|
|
109
|
+
async def run_continue(self, *, ctx):
|
|
110
|
+
return await self.runs.continue_run(ack=ctx.ack)
|
|
97
111
|
|
|
98
112
|
@command("run/input")
|
|
99
113
|
async def run_input(self, params):
|
|
@@ -27,6 +27,7 @@ async def test_multi_steer_batch_places_all_and_dispatches_once():
|
|
|
27
27
|
assert [m["id"] for m in call.ctx.messages] == ["s1", "s2", "s3"]
|
|
28
28
|
script.no_call()
|
|
29
29
|
assert queue_ids(drv.replica, "steerQueue") == []
|
|
30
|
+
call.ack()
|
|
30
31
|
call.finish(RunManager.Complete())
|
|
31
32
|
for offset in range(3):
|
|
32
33
|
assert (await drv.res(first + offset))["type"] == "accepted"
|
|
@@ -40,7 +41,7 @@ async def test_steer_and_stop_in_one_batch_nets_to_stop():
|
|
|
40
41
|
first = await drv.batch(
|
|
41
42
|
[("run/steer", add("s1")), ("run/stop", None)]
|
|
42
43
|
)
|
|
43
|
-
assert (await drv.res(first
|
|
44
|
+
assert (await drv.res(first))["type"] == "accepted"
|
|
44
45
|
assert (await drv.res(first + 1, terminal=False))["type"] == "pending"
|
|
45
46
|
await asyncio.wait_for(call.ctx.stop_requested.wait(), 5)
|
|
46
47
|
call.finish(RunManager.Stop(dispatch_queue=False))
|
|
@@ -115,7 +115,7 @@ async def test_anchor_on_the_queued_edit_form_is_validated_never_ignored():
|
|
|
115
115
|
assert drv.replica["queue"][0]["parts"][0]["text"] == "edited"
|
|
116
116
|
|
|
117
117
|
|
|
118
|
-
async def
|
|
118
|
+
async def test_post_reload_sends_anchor_on_the_reload_targets_parent():
|
|
119
119
|
script = Script()
|
|
120
120
|
script.thread.update(
|
|
121
121
|
{
|
|
@@ -124,22 +124,14 @@ async def test_reload_requires_a_message_id_and_delivers_it_on_the_start_context
|
|
|
124
124
|
}
|
|
125
125
|
)
|
|
126
126
|
async with run_host(script, capabilities=("rewind",)) as (drv, host):
|
|
127
|
-
|
|
128
|
-
await drv.command("run/reload", {"sourceId": "a1"}, terminal=False),
|
|
129
|
-
"invalid-message",
|
|
130
|
-
)
|
|
131
|
-
assert_rejected(
|
|
132
|
-
await drv.command(
|
|
133
|
-
"run/reload", {"sourceId": "a1", "messageId": "u1"}, terminal=False
|
|
134
|
-
),
|
|
135
|
-
"duplicate-id",
|
|
136
|
-
)
|
|
137
|
-
res = await drv.command(
|
|
138
|
-
"run/reload", {"sourceId": "a1", "messageId": "r9"}, terminal=False
|
|
139
|
-
)
|
|
127
|
+
res = await drv.command("run/reload", {"sourceId": "a1"}, terminal=False)
|
|
140
128
|
assert res["type"] == "pending"
|
|
141
129
|
call = await script.next_call()
|
|
142
130
|
assert call.ctx.type == "message-reload"
|
|
143
|
-
|
|
131
|
+
script.thread["a1"] = {
|
|
132
|
+
"parentId": "u1", "role": "assistant", "isLeaf": True, "onActiveBranch": False
|
|
133
|
+
}
|
|
134
|
+
res = await drv.command("run/enqueue", add("m1", anchor="u1"), terminal=False)
|
|
135
|
+
assert res["type"] == "pending"
|
|
144
136
|
call.finish(RunManager.Complete())
|
|
145
|
-
await drv.wait_status("
|
|
137
|
+
await drv.wait_status("running")
|
|
@@ -13,6 +13,7 @@ async def test_immediate_dispatch_carries_caller():
|
|
|
13
13
|
call = await script.next_call()
|
|
14
14
|
assert call.ctx.caller is not None
|
|
15
15
|
assert call.ctx.caller.client_id == "c1"
|
|
16
|
+
call.ack()
|
|
16
17
|
call.finish(RunManager.Complete())
|
|
17
18
|
|
|
18
19
|
|
|
@@ -24,6 +25,7 @@ async def test_queue_entry_projects_caller_client_id():
|
|
|
24
25
|
await drv.command("run/enqueue", add("m2"), terminal=False)
|
|
25
26
|
await drv.wait(lambda s: len(s["queue"]) == 1)
|
|
26
27
|
assert drv.replica["queue"][0]["caller"] == {"clientId": "c1"}
|
|
28
|
+
call.ack()
|
|
27
29
|
call.finish(RunManager.Complete())
|
|
28
30
|
|
|
29
31
|
|
|
@@ -34,12 +36,14 @@ async def test_caller_context_is_read_at_dispatch_time():
|
|
|
34
36
|
call = await script.next_call()
|
|
35
37
|
await drv.command("run/enqueue", add("m2"), terminal=False)
|
|
36
38
|
assert (await drv.context({"user": "simon"}))["type"] == "accepted"
|
|
39
|
+
call.ack()
|
|
37
40
|
call.finish(RunManager.Complete())
|
|
38
41
|
dispatched = await script.next_call()
|
|
39
42
|
assert dispatched.ctx.caller is not None
|
|
40
43
|
assert dispatched.ctx.caller.context == {"user": "simon"}
|
|
41
44
|
assert dispatched.ctx.messages[0]["id"] == "m2"
|
|
42
45
|
assert "caller" not in dispatched.ctx.messages[0]
|
|
46
|
+
dispatched.ack()
|
|
43
47
|
dispatched.finish(RunManager.Complete())
|
|
44
48
|
|
|
45
49
|
|
|
@@ -47,11 +51,12 @@ async def test_dispatch_without_queue_entry_has_no_caller():
|
|
|
47
51
|
script = Script()
|
|
48
52
|
async with run_host(script, capabilities=("rewind",)) as (drv, _):
|
|
49
53
|
script.thread["a1"] = {"role": "assistant", "parentId": None, "isLeaf": True}
|
|
50
|
-
assert (await drv.command("run/reload", {"sourceId": "a1"
|
|
54
|
+
assert (await drv.command("run/reload", {"sourceId": "a1"}, terminal=False))[
|
|
51
55
|
"type"
|
|
52
56
|
] == "pending"
|
|
53
57
|
call = await script.next_call()
|
|
54
58
|
assert call.ctx.caller is None
|
|
59
|
+
call.ack()
|
|
55
60
|
call.finish(RunManager.Complete())
|
|
56
61
|
|
|
57
62
|
|
|
@@ -65,6 +70,7 @@ async def test_take_steered_strips_the_caller_stamp():
|
|
|
65
70
|
assert drv.replica["steerQueue"][0]["caller"] == {"clientId": "c1"}
|
|
66
71
|
[taken] = call.ctx.take_steered()
|
|
67
72
|
assert "caller" not in taken
|
|
73
|
+
call.ack()
|
|
68
74
|
call.finish(RunManager.Complete())
|
|
69
75
|
|
|
70
76
|
|
|
@@ -85,6 +91,7 @@ async def test_max_queued_caps_each_lane():
|
|
|
85
91
|
] == "pending"
|
|
86
92
|
steer_full = await drv.command("run/steer", add("s2"), terminal=False)
|
|
87
93
|
assert steer_full["payload"] == {"reason": "queue-full"}
|
|
94
|
+
call.ack()
|
|
88
95
|
call.finish(RunManager.Complete())
|
|
89
96
|
|
|
90
97
|
|
|
@@ -97,6 +104,7 @@ async def test_lane_change_into_a_full_lane_rejects():
|
|
|
97
104
|
await drv.command("run/steer", add("s1"), terminal=False)
|
|
98
105
|
moved = await drv.command("run/steer", {"messageId": "m2"}, terminal=False)
|
|
99
106
|
assert moved["payload"] == {"reason": "queue-full"}
|
|
107
|
+
call.ack()
|
|
100
108
|
call.finish(RunManager.Complete())
|
|
101
109
|
|
|
102
110
|
|
|
@@ -109,5 +117,6 @@ async def test_max_queued_below_one_rejects_at_construction():
|
|
|
109
117
|
get_message_meta=script.get_message_meta,
|
|
110
118
|
create_task=lambda coro: None,
|
|
111
119
|
schedule=lambda fn: None,
|
|
120
|
+
leaf_message_id=None,
|
|
112
121
|
max_queued=0,
|
|
113
122
|
)
|