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.
Files changed (25) hide show
  1. {harness_sdk_python-0.4.0 → harness_sdk_python-0.4.2}/PKG-INFO +1 -1
  2. {harness_sdk_python-0.4.0 → harness_sdk_python-0.4.2}/pyproject.toml +1 -1
  3. {harness_sdk_python-0.4.0 → harness_sdk_python-0.4.2}/src/harness_sdk/run_manager.py +126 -70
  4. {harness_sdk_python-0.4.0 → harness_sdk_python-0.4.2}/tests/run_helpers.py +16 -2
  5. {harness_sdk_python-0.4.0 → harness_sdk_python-0.4.2}/tests/test_batches.py +2 -1
  6. {harness_sdk_python-0.4.0 → harness_sdk_python-0.4.2}/tests/test_branch_anchor.py +8 -16
  7. {harness_sdk_python-0.4.0 → harness_sdk_python-0.4.2}/tests/test_caller.py +10 -1
  8. {harness_sdk_python-0.4.0 → harness_sdk_python-0.4.2}/tests/test_edit_dispatched.py +5 -0
  9. {harness_sdk_python-0.4.0 → harness_sdk_python-0.4.2}/tests/test_edit_reload.py +13 -9
  10. {harness_sdk_python-0.4.0 → harness_sdk_python-0.4.2}/tests/test_enqueue.py +13 -1
  11. {harness_sdk_python-0.4.0 → harness_sdk_python-0.4.2}/tests/test_facade.py +6 -3
  12. {harness_sdk_python-0.4.0 → harness_sdk_python-0.4.2}/tests/test_input_required.py +11 -2
  13. {harness_sdk_python-0.4.0 → harness_sdk_python-0.4.2}/tests/test_outcomes.py +13 -2
  14. {harness_sdk_python-0.4.0 → harness_sdk_python-0.4.2}/tests/test_rewind_during_run.py +9 -4
  15. harness_sdk_python-0.4.2/tests/test_run_leaf.py +121 -0
  16. harness_sdk_python-0.4.2/tests/test_settle.py +218 -0
  17. {harness_sdk_python-0.4.0 → harness_sdk_python-0.4.2}/tests/test_steer.py +5 -0
  18. {harness_sdk_python-0.4.0 → harness_sdk_python-0.4.2}/tests/test_stop_continue.py +32 -5
  19. harness_sdk_python-0.4.0/tests/test_settle.py +0 -131
  20. {harness_sdk_python-0.4.0 → harness_sdk_python-0.4.2}/.gitignore +0 -0
  21. {harness_sdk_python-0.4.0 → harness_sdk_python-0.4.2}/README.md +0 -0
  22. {harness_sdk_python-0.4.0 → harness_sdk_python-0.4.2}/src/harness_sdk/__init__.py +0 -0
  23. {harness_sdk_python-0.4.0 → harness_sdk_python-0.4.2}/src/harness_sdk/fenced_postgres.py +0 -0
  24. {harness_sdk_python-0.4.0 → harness_sdk_python-0.4.2}/tests/test_fenced_postgres.py +0 -0
  25. {harness_sdk_python-0.4.0 → harness_sdk_python-0.4.2}/tests/test_placement.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: harness-sdk-python
3
- Version: 0.4.0
3
+ Version: 0.4.2
4
4
  Summary: RunManager: the harness-sdk runs subsystem for Python Statewire hosts
5
5
  Project-URL: Repository, https://github.com/assistant-ui/harness-sdk
6
6
  License-Expression: MIT
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "harness-sdk-python"
3
- version = "0.4.0"
3
+ version = "0.4.2"
4
4
  description = "RunManager: the harness-sdk runs subsystem for Python Statewire hosts"
5
5
  readme = "README.md"
6
6
  license = "MIT"
@@ -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). Late-settling commands (rewinds, stop) await a future the drain
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
- pass
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._entity_futures: dict[str, "asyncio.Future[Any]"] = {}
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
- self._reload_message_id = message_id
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._settle_entities(exc)
419
+ self._settle_initiators(exc)
386
420
  else:
387
421
  self._freeze(message)
388
- self._settle_entities(_reject("run-error", message))
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
- self._settle_entities(
393
- _reject("run-error", "run ended in error")
394
- if isinstance(outcome, RunManager.Error)
395
- else None
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 _settle_entities(self, error: StatewireReject | None) -> None:
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 _adopt_entities(self, ids: Iterable[str]) -> None:
417
- for id in ids:
418
- future = self._entity_futures.pop(id, None)
419
- if future is not None:
420
- self._run_futures.append(future)
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 or e.anchor == self._reload_message_id:
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 or e.message_id == self._reload_message_id:
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
- self._entity_futures[e.message_id] = e.future
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, message_id, ack))
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) -> None:
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
- self._intake.append(_Continue())
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
- self._manager._callers.pop(item["id"], None)
1103
- self._manager._adopt_entities(item["id"] for item in items)
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, terminal=False))["type"] == "pending"
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 test_reload_requires_a_message_id_and_delivers_it_on_the_start_context():
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
- assert_rejected(
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
- assert call.ctx.message_id == "r9"
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("ready")
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", "messageId": "r1"}, terminal=False))[
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
  )