harness-sdk-python 0.5.0__tar.gz → 0.7.0__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 (31) hide show
  1. {harness_sdk_python-0.5.0 → harness_sdk_python-0.7.0}/.gitignore +1 -0
  2. {harness_sdk_python-0.5.0 → harness_sdk_python-0.7.0}/PKG-INFO +2 -2
  3. {harness_sdk_python-0.5.0 → harness_sdk_python-0.7.0}/README.md +1 -1
  4. {harness_sdk_python-0.5.0 → harness_sdk_python-0.7.0}/pyproject.toml +1 -1
  5. harness_sdk_python-0.7.0/src/harness_sdk/__init__.py +4 -0
  6. {harness_sdk_python-0.5.0 → harness_sdk_python-0.7.0}/src/harness_sdk/fenced_postgres.py +58 -18
  7. harness_sdk_python-0.7.0/src/harness_sdk/linear_thread.py +52 -0
  8. {harness_sdk_python-0.5.0 → harness_sdk_python-0.7.0}/src/harness_sdk/run_manager.py +79 -34
  9. {harness_sdk_python-0.5.0 → harness_sdk_python-0.7.0}/tests/run_helpers.py +72 -14
  10. harness_sdk_python-0.7.0/tests/test_ack_visibility.py +83 -0
  11. {harness_sdk_python-0.5.0 → harness_sdk_python-0.7.0}/tests/test_batches.py +12 -11
  12. {harness_sdk_python-0.5.0 → harness_sdk_python-0.7.0}/tests/test_branch_anchor.py +12 -17
  13. {harness_sdk_python-0.5.0 → harness_sdk_python-0.7.0}/tests/test_edit_dispatched.py +14 -23
  14. {harness_sdk_python-0.5.0 → harness_sdk_python-0.7.0}/tests/test_edit_reload.py +19 -28
  15. {harness_sdk_python-0.5.0 → harness_sdk_python-0.7.0}/tests/test_enqueue.py +28 -30
  16. {harness_sdk_python-0.5.0 → harness_sdk_python-0.7.0}/tests/test_facade.py +1 -1
  17. {harness_sdk_python-0.5.0 → harness_sdk_python-0.7.0}/tests/test_fenced_postgres.py +64 -0
  18. {harness_sdk_python-0.5.0 → harness_sdk_python-0.7.0}/tests/test_input_required.py +16 -18
  19. harness_sdk_python-0.7.0/tests/test_linear_thread.py +122 -0
  20. {harness_sdk_python-0.5.0 → harness_sdk_python-0.7.0}/tests/test_meta.py +16 -22
  21. {harness_sdk_python-0.5.0 → harness_sdk_python-0.7.0}/tests/test_outcomes.py +15 -14
  22. {harness_sdk_python-0.5.0 → harness_sdk_python-0.7.0}/tests/test_placement.py +2 -2
  23. harness_sdk_python-0.7.0/tests/test_prepare_hooks.py +164 -0
  24. {harness_sdk_python-0.5.0 → harness_sdk_python-0.7.0}/tests/test_rewind_during_run.py +14 -28
  25. {harness_sdk_python-0.5.0 → harness_sdk_python-0.7.0}/tests/test_run_leaf.py +53 -48
  26. {harness_sdk_python-0.5.0 → harness_sdk_python-0.7.0}/tests/test_settle.py +36 -60
  27. {harness_sdk_python-0.5.0 → harness_sdk_python-0.7.0}/tests/test_steer.py +16 -12
  28. {harness_sdk_python-0.5.0 → harness_sdk_python-0.7.0}/tests/test_stop_continue.py +28 -24
  29. harness_sdk_python-0.5.0/src/harness_sdk/__init__.py +0 -4
  30. harness_sdk_python-0.5.0/src/harness_sdk/linear_thread.py +0 -31
  31. harness_sdk_python-0.5.0/tests/test_linear_thread.py +0 -87
@@ -20,3 +20,4 @@ __pycache__
20
20
  # agentdoc mounted docs
21
21
  /d[0-9]*.md
22
22
  /doc_*.md
23
+ apps/docs/.docs
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: harness-sdk-python
3
- Version: 0.5.0
3
+ Version: 0.7.0
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
@@ -23,7 +23,7 @@ class MyHost(Statewire):
23
23
  self.runs = RunManager(
24
24
  state=self.state,
25
25
  start=self._start,
26
- get_message_meta=self._get_message_meta,
26
+ thread=self._thread,
27
27
  create_task=self.create_task,
28
28
  capabilities=("rewind",),
29
29
  )
@@ -11,7 +11,7 @@ class MyHost(Statewire):
11
11
  self.runs = RunManager(
12
12
  state=self.state,
13
13
  start=self._start,
14
- get_message_meta=self._get_message_meta,
14
+ thread=self._thread,
15
15
  create_task=self.create_task,
16
16
  capabilities=("rewind",),
17
17
  )
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "harness-sdk-python"
3
- version = "0.5.0"
3
+ version = "0.7.0"
4
4
  description = "RunManager: the harness-sdk runs subsystem for Python Statewire hosts"
5
5
  readme = "README.md"
6
6
  license = "MIT"
@@ -0,0 +1,4 @@
1
+ from .linear_thread import linear_thread
2
+ from .run_manager import RunManager
3
+
4
+ __all__ = ["RunManager", "linear_thread"]
@@ -7,7 +7,7 @@ try:
7
7
  from langchain_core.runnables import RunnableConfig
8
8
  from langgraph.checkpoint.base import SerializerProtocol
9
9
  from langgraph.checkpoint.postgres.aio import AsyncPostgresSaver, _ainternal
10
- from psycopg import AsyncPipeline, errors
10
+ from psycopg import AsyncConnection, AsyncPipeline, errors
11
11
  from psycopg.pq import TransactionStatus
12
12
  from psycopg.cursor_async import AsyncCursor
13
13
  from psycopg.rows import DictRow, dict_row
@@ -126,6 +126,57 @@ class AsyncFencedPostgresSaver(AsyncPostgresSaver):
126
126
  self.on_fence_lost = on_fence_lost
127
127
  self._fence_lost_notified = False
128
128
 
129
+ @staticmethod
130
+ async def _validate_and_bump(
131
+ cur: "AsyncCursor[DictRow]",
132
+ thread_id: str,
133
+ *,
134
+ create_fence_table: bool,
135
+ fence_table: str,
136
+ ) -> int:
137
+ if create_fence_table:
138
+ await cur.execute(_create_sql(fence_table))
139
+ await cur.execute(_create_fn_sql(fence_table))
140
+ else:
141
+ await cur.execute(
142
+ "SELECT to_regclass(%s) AS tbl, to_regprocedure(%s) AS fn",
143
+ (fence_table, f"{_assert_fn(fence_table)}(text, bigint)"),
144
+ )
145
+ row = await cur.fetchone()
146
+ if row["tbl"] is None:
147
+ raise _missing_fence_error(f"fence table {fence_table!r}")
148
+ if row["fn"] is None:
149
+ raise _missing_fence_error(
150
+ f"fence function {_assert_fn(fence_table)!r}"
151
+ )
152
+ await cur.execute(_bump_sql(fence_table), (thread_id,))
153
+ return (await cur.fetchone())["epoch"]
154
+
155
+ @classmethod
156
+ async def bump_epoch(
157
+ cls,
158
+ conn: "AsyncConnection[Any]",
159
+ thread_id: str,
160
+ *,
161
+ create_fence_table: bool,
162
+ fence_table: str = DEFAULT_FENCE_TABLE,
163
+ ) -> int:
164
+ """Bump ``thread_id``'s epoch inside the caller's open transaction; the epoch is effective only once the caller commits, which must happen before it is used for writes."""
165
+ _require_thread_id(thread_id)
166
+ _require_fence_table(fence_table)
167
+ if conn.info.transaction_status != TransactionStatus.INTRANS:
168
+ raise RuntimeError(
169
+ "bump_epoch requires a connection with a transaction in "
170
+ "progress; use acquire for idle connections"
171
+ )
172
+ async with conn.cursor(binary=True, row_factory=dict_row) as cur:
173
+ return await cls._validate_and_bump(
174
+ cur,
175
+ thread_id,
176
+ create_fence_table=create_fence_table,
177
+ fence_table=fence_table,
178
+ )
179
+
129
180
  @classmethod
130
181
  async def acquire(
131
182
  cls,
@@ -149,23 +200,12 @@ class AsyncFencedPostgresSaver(AsyncPostgresSaver):
149
200
  c.transaction(),
150
201
  c.cursor(binary=True, row_factory=dict_row) as cur,
151
202
  ):
152
- if create_fence_table:
153
- await cur.execute(_create_sql(fence_table))
154
- await cur.execute(_create_fn_sql(fence_table))
155
- else:
156
- await cur.execute(
157
- "SELECT to_regclass(%s) AS tbl, to_regprocedure(%s) AS fn",
158
- (fence_table, f"{_assert_fn(fence_table)}(text, bigint)"),
159
- )
160
- row = await cur.fetchone()
161
- if row["tbl"] is None:
162
- raise _missing_fence_error(f"fence table {fence_table!r}")
163
- if row["fn"] is None:
164
- raise _missing_fence_error(
165
- f"fence function {_assert_fn(fence_table)!r}"
166
- )
167
- await cur.execute(_bump_sql(fence_table), (thread_id,))
168
- epoch = (await cur.fetchone())["epoch"]
203
+ epoch = await cls._validate_and_bump(
204
+ cur,
205
+ thread_id,
206
+ create_fence_table=create_fence_table,
207
+ fence_table=fence_table,
208
+ )
169
209
  if pipe is not None:
170
210
  await pipe.sync()
171
211
  return cls(
@@ -0,0 +1,52 @@
1
+ from dataclasses import dataclass
2
+ from typing import Any, Callable
3
+
4
+
5
+ @dataclass(frozen=True)
6
+ class LinearThread:
7
+ _messages: Callable[[], list[dict[str, Any]]]
8
+ _role: Callable[[dict[str, Any]], str]
9
+
10
+ async def get_message_meta(self, message_id: str | None) -> dict[str, Any] | None:
11
+ items = self._messages()
12
+ if message_id is None:
13
+ return {"isLeaf": not items}
14
+ for index, message in enumerate(items):
15
+ if message.get("id") == message_id:
16
+ return {
17
+ "parentId": items[index - 1].get("id") if index > 0 else None,
18
+ "role": self._role(message),
19
+ "isLeaf": index == len(items) - 1,
20
+ "onActiveBranch": True,
21
+ }
22
+ return None
23
+
24
+ async def get_leaf_message_id(self) -> str | None:
25
+ items = self._messages()
26
+ return items[-1].get("id") if items else None
27
+
28
+ async def get_message_child_id(self, parent_id: str) -> str | None:
29
+ """The id of the first non-tool message after parent_id; None when parent_id is unknown or only tools follow."""
30
+ items = self._messages()
31
+ index = next(
32
+ (i for i, m in enumerate(items) if m.get("id") == parent_id), None
33
+ )
34
+ if index is None:
35
+ return None
36
+ for message in items[index + 1 :]:
37
+ if message.get("type") != "tool":
38
+ return message.get("id")
39
+ return None
40
+
41
+
42
+ def linear_thread(
43
+ *,
44
+ messages: Callable[[], list[dict[str, Any]]],
45
+ role: Callable[[dict[str, Any]], str],
46
+ ) -> LinearThread:
47
+ """RunManager thread projection over a linear message list: ``messages`` returns the current list, ``role`` maps a message to its role."""
48
+ if not callable(messages):
49
+ raise TypeError("messages must be callable")
50
+ if not callable(role):
51
+ raise TypeError("role must be callable")
52
+ return LinearThread(messages, role)
@@ -19,7 +19,7 @@ initiators settle rejected (``stopped`` or the failure).
19
19
  import asyncio
20
20
  import uuid
21
21
  from dataclasses import dataclass, field
22
- from typing import Any, Awaitable, Callable, Iterable
22
+ from typing import Any, Awaitable, Callable, Iterable, Protocol
23
23
 
24
24
  from statewire import StatewireReject
25
25
  from statewire.state import plain
@@ -51,10 +51,6 @@ _TRIGGERS = (
51
51
 
52
52
  _DECISIONS = ("approve", "reject", "edit", "respond")
53
53
 
54
- # Called with None it answers for the thread root: isLeaf True iff the thread is empty.
55
- GetMessageMeta = Callable[[str | None], Awaitable[dict[str, Any] | None]]
56
-
57
-
58
54
  def _reject(reason: str, message: str) -> StatewireReject:
59
55
  return StatewireReject(message, payload={"reason": reason})
60
56
 
@@ -129,7 +125,6 @@ class _Rewind:
129
125
  ack: Callable[[], None]
130
126
  future: "asyncio.Future[Any]"
131
127
  root_meta: Any = None
132
- acked: bool = False
133
128
 
134
129
 
135
130
  @dataclass
@@ -138,21 +133,38 @@ class _Effects:
138
133
  new_added: bool = False
139
134
  continue_requested: bool = False
140
135
  continue_meta: Any = None
141
- staged_sends: list[tuple[str, "asyncio.Future[Any]"]] = field(default_factory=list)
136
+ staged_sends: list[tuple[str, "asyncio.Future[Any]", Callable[[], None]]] = field(
137
+ default_factory=list
138
+ )
142
139
  continues: list["asyncio.Future[Any]"] = field(default_factory=list)
143
140
 
144
141
 
145
142
  class RunManager:
143
+ class Thread(Protocol):
144
+ """The thread projection RunManager reads; the host owns the tree."""
145
+
146
+ async def get_message_meta(
147
+ self, message_id: str | None
148
+ ) -> dict[str, Any] | None:
149
+ """Meta for a known id ({parentId, role, isLeaf, onActiveBranch}), None for an unknown one; a None id probes the root ({isLeaf})."""
150
+ ...
151
+
152
+ async def get_leaf_message_id(self) -> str | None:
153
+ """The active branch's current leaf id, None while the thread is empty."""
154
+ ...
155
+
146
156
  def __init__(
147
157
  self,
148
158
  *,
149
159
  state: Any,
150
160
  run: Callable[["RunManager.RunContext"], Awaitable[Any]],
151
- get_message_meta: GetMessageMeta,
161
+ thread: "RunManager.Thread",
152
162
  create_task: Callable[[Any], "asyncio.Task[Any]"],
153
163
  schedule: Callable[[Callable[[], None]], None],
154
164
  capabilities: Iterable[str] = (),
155
165
  max_queued: int = 50,
166
+ prepare_message: Callable[[dict[str, Any]], dict[str, Any]] | None = None,
167
+ prepare_input: Callable[[dict[str, Any]], dict[str, Any]] | None = None,
156
168
  ) -> None:
157
169
  caps = frozenset(capabilities)
158
170
  unknown = caps - _CAPABILITIES
@@ -164,11 +176,13 @@ class RunManager:
164
176
  raise ValueError("max_queued must be >= 1")
165
177
  self._state = state
166
178
  self._run = run
167
- self._get_message_meta = get_message_meta
179
+ self._thread = thread
168
180
  self._capabilities = caps
169
181
  self._create_task = create_task
170
182
  self._schedule = schedule
171
183
  self._max_queued = max_queued
184
+ self._prepare_message = prepare_message
185
+ self._prepare_input = prepare_input
172
186
  self._task: "asyncio.Task[Any] | None" = None
173
187
  self._ctx: "RunManager.RunContext | None" = None
174
188
  self._dispatched_ids: tuple[str, ...] = ()
@@ -179,6 +193,10 @@ class RunManager:
179
193
  self._stop_reason: str | None = None
180
194
  self._staged_rewinds: list[_Rewind] = []
181
195
  self._dispatching: list[tuple[str, dict[str, Any]]] = []
196
+ # Deferred covering acks: they only ride an envelope whose state shows
197
+ # the command's effect (message-ack append, or the revert re-insert).
198
+ self._deferred_acks: dict[str, Callable[[], None]] = {}
199
+ self._dispatch_ack: Callable[[], None] | None = None
182
200
  self._run_acked = True
183
201
  self._leaf_confirmed: str | None = None
184
202
  self._send_futures: list["asyncio.Future[Any]"] = []
@@ -294,7 +312,7 @@ class RunManager:
294
312
  rollback_to=rewind.rollback_to,
295
313
  root_meta=rewind.root_meta,
296
314
  )
297
- rewind.ack()
315
+ self._dispatch_ack = rewind.ack
298
316
  if not rewind.future.done():
299
317
  self._run_futures.append(rewind.future)
300
318
  elif outcome is not None:
@@ -304,16 +322,15 @@ class RunManager:
304
322
  if self._task is not None and (self._staged_stops or self._staged_rewinds):
305
323
  assert self._ctx is not None
306
324
  self._ctx.stop_requested.set()
307
- for rewind in self._staged_rewinds:
308
- if not rewind.acked:
309
- rewind.ack()
310
- rewind.acked = True
311
325
  dispatching = {item["id"] for _, item in self._dispatching}
312
- for message_id, future in fx.staged_sends:
326
+ for message_id, future, ack in fx.staged_sends:
313
327
  if message_id in dispatching:
328
+ self._deferred_acks[message_id] = ack
314
329
  self._send_futures.append(future)
315
- elif not future.done():
316
- future.set_result(None)
330
+ else:
331
+ ack()
332
+ if not future.done():
333
+ future.set_result(None)
317
334
  for future in fx.continues:
318
335
  if not future.done():
319
336
  future.set_result(None)
@@ -480,6 +497,7 @@ class RunManager:
480
497
  "run settled without acking its messages (call ctx.ack_messages())"
481
498
  )
482
499
  except asyncio.CancelledError:
500
+ await self._pull_leaf()
483
501
  self._settle(ctx)
484
502
  self._set_status("stopped")
485
503
  self._entry()["runId"] = None
@@ -488,6 +506,7 @@ class RunManager:
488
506
  self._idle.set()
489
507
  raise # no drain, no freeze
490
508
  except Exception as exc:
509
+ await self._pull_leaf()
491
510
  self._settle(ctx)
492
511
  message = str(exc) or type(exc).__name__
493
512
  if isinstance(exc, StatewireReject):
@@ -499,6 +518,7 @@ class RunManager:
499
518
  self._revert_dispatching()
500
519
  self._drain()
501
520
  return
521
+ await self._pull_leaf()
502
522
  self._settle(ctx)
503
523
  if isinstance(outcome, RunManager.Error):
504
524
  error = _reject("run-error", "run ended in error")
@@ -511,11 +531,18 @@ class RunManager:
511
531
  self._outcome = outcome
512
532
  self._drain()
513
533
 
534
+ async def _pull_leaf(self) -> None:
535
+ # An unacked end reverts instead of recording a leaf.
536
+ if self._dispatching or not self._run_acked:
537
+ return
538
+ self._entry()["runLeafMessageId"] = await self._thread.get_leaf_message_id()
539
+
514
540
  def _settle(self, ctx: "RunManager.RunContext") -> None:
515
541
  if self._ctx is ctx:
516
542
  self._ctx = None
517
543
  self._task = None
518
544
  self._dispatch_record = None
545
+ self._dispatch_ack = None
519
546
 
520
547
  def _settle_initiators(self, error: StatewireReject | None) -> None:
521
548
  futures, self._run_futures = self._run_futures, []
@@ -534,6 +561,13 @@ class RunManager:
534
561
  future.set_result(None)
535
562
 
536
563
  def _ack_messages(self) -> None:
564
+ if self._dispatch_ack is not None:
565
+ self._dispatch_ack()
566
+ self._dispatch_ack = None
567
+ for _, item in self._dispatching:
568
+ ack = self._deferred_acks.pop(item["id"], None)
569
+ if ack is not None:
570
+ ack()
537
571
  self._dispatching = []
538
572
  self._run_acked = True
539
573
  self._leaf_confirmed = plain(self._entry()["runLeafMessageId"])
@@ -550,6 +584,10 @@ class RunManager:
550
584
  front = [item for taken_lane, item in taken if taken_lane == lane]
551
585
  if front:
552
586
  entry[lane] = front + self._lane_items(lane)
587
+ for _, item in taken:
588
+ ack = self._deferred_acks.pop(item["id"], None)
589
+ if ack is not None:
590
+ ack()
553
591
  ids = {item["id"] for _, item in taken}
554
592
  self._dispatched_ids = tuple(id for id in self._dispatched_ids if id not in ids)
555
593
  if plain(entry["runLeafMessageId"]) in ids:
@@ -585,7 +623,22 @@ class RunManager:
585
623
 
586
624
  # ─── Message and placement validation ───────────────────
587
625
 
626
+ def _prepared(
627
+ self,
628
+ hook: Callable[[dict[str, Any]], dict[str, Any]] | None,
629
+ name: str,
630
+ value: Any,
631
+ ) -> Any:
632
+ # Non-dict input skips the hook and falls through to validation's reject.
633
+ if hook is None or not isinstance(value, dict):
634
+ return value
635
+ result = hook(value)
636
+ if not isinstance(result, dict):
637
+ raise _reject("invalid-message", f"{name} must return an object")
638
+ return result
639
+
588
640
  def _validated_message(self, message: Any) -> dict[str, Any]:
641
+ message = self._prepared(self._prepare_message, "prepare_message", message)
589
642
  if not isinstance(message, dict):
590
643
  raise _reject("invalid-message", "message must be an object")
591
644
  if not isinstance(message.get("id"), str) or message["id"] == "":
@@ -772,8 +825,7 @@ class RunManager:
772
825
  raise _reject("invalid-message", "anchorMessageId is required")
773
826
  self._check_anchor(e)
774
827
  self._insert_new(e.lane, e.message, e.params, e.meta)
775
- fx.staged_sends.append((e.message_id, e.future))
776
- e.ack()
828
+ fx.staged_sends.append((e.message_id, e.future, e.ack))
777
829
  fx.new_added = True
778
830
  if e.lane == "steerQueue":
779
831
  fx.steer_added = True
@@ -907,7 +959,7 @@ class RunManager:
907
959
  anchor_meta: dict[str, Any] | None = None
908
960
  thread_empty = False
909
961
  if anchor is None:
910
- root = await self._get_message_meta(None)
962
+ root = await self._thread.get_message_meta(None)
911
963
  assert root is not None, "get_message_meta(None) must answer for the root"
912
964
  thread_empty = bool(root["isLeaf"])
913
965
  elif anchor is not _ABSENT:
@@ -915,7 +967,7 @@ class RunManager:
915
967
  raise _reject(
916
968
  "invalid-message", "anchorMessageId must be an id or null"
917
969
  )
918
- anchor_meta = await self._get_message_meta(anchor)
970
+ anchor_meta = await self._thread.get_message_meta(anchor)
919
971
  if not has_message:
920
972
  message_id = params["messageId"]
921
973
  if not isinstance(message_id, str):
@@ -935,7 +987,7 @@ class RunManager:
935
987
  )
936
988
  )
937
989
  message = self._validated_message(params["message"])
938
- source_meta = await self._get_message_meta(message["id"])
990
+ source_meta = await self._thread.get_message_meta(message["id"])
939
991
  return await self._stage(
940
992
  _Send(
941
993
  lane,
@@ -984,7 +1036,7 @@ class RunManager:
984
1036
  source_id = params.get("sourceId") if isinstance(params, dict) else None
985
1037
  if not isinstance(source_id, str):
986
1038
  raise _reject("invalid-message", "sourceId must be a string")
987
- source_meta = await self._get_message_meta(source_id)
1039
+ source_meta = await self._thread.get_message_meta(source_id)
988
1040
  if source_meta is None:
989
1041
  raise _reject("unknown-id", f"message {source_id} is unknown")
990
1042
  if source_meta["role"] != "user" and "assistant-edit" not in self._capabilities:
@@ -996,7 +1048,7 @@ class RunManager:
996
1048
  )
997
1049
  if (
998
1050
  message["id"] != source_id
999
- and await self._get_message_meta(message["id"]) is not None
1051
+ and await self._thread.get_message_meta(message["id"]) is not None
1000
1052
  ):
1001
1053
  raise _reject("duplicate-id", f"message id {message['id']} is already used")
1002
1054
  return await self._stage(_Edit(source_id, source_meta, message, meta, ack))
@@ -1008,13 +1060,13 @@ class RunManager:
1008
1060
  source_id = params.get("sourceId") if isinstance(params, dict) else None
1009
1061
  if not isinstance(source_id, str):
1010
1062
  raise _reject("invalid-message", "sourceId must be a string")
1011
- source_meta = await self._get_message_meta(source_id)
1063
+ source_meta = await self._thread.get_message_meta(source_id)
1012
1064
  if source_meta is None:
1013
1065
  raise _reject("unknown-id", f"message {source_id} is unknown")
1014
1066
  if source_meta["role"] != "assistant":
1015
1067
  raise _reject("invalid-message", "sourceId must name an assistant message")
1016
1068
  if source_meta["parentId"] is not None:
1017
- parent = await self._get_message_meta(source_meta["parentId"])
1069
+ parent = await self._thread.get_message_meta(source_meta["parentId"])
1018
1070
  if (
1019
1071
  parent is not None
1020
1072
  and parent["role"] == "assistant"
@@ -1049,6 +1101,7 @@ class RunManager:
1049
1101
  return await entry.future
1050
1102
 
1051
1103
  def _validated_response(self, request_type: str, response: Any) -> dict[str, Any]:
1104
+ response = self._prepared(self._prepare_input, "prepare_input", response)
1052
1105
  if not isinstance(response, dict):
1053
1106
  raise _reject("invalid-message", "response must be an object")
1054
1107
  if request_type == "tool-call":
@@ -1230,14 +1283,6 @@ class RunManager:
1230
1283
  self._ensure_active()
1231
1284
  self._manager._ack_messages()
1232
1285
 
1233
- def set_leaf_message_id(self, message_id: str) -> None:
1234
- self._ensure_active()
1235
- if not isinstance(message_id, str) or message_id == "":
1236
- raise ValueError("message_id must be a non-empty string")
1237
- if self._manager._dispatching or not self._manager._run_acked:
1238
- raise RuntimeError("ack_messages must precede set_leaf_message_id")
1239
- self._manager._entry()["runLeafMessageId"] = message_id
1240
-
1241
1286
  def set_recovery_state(self, value: Any) -> None:
1242
1287
  self._ensure_active()
1243
1288
  record = self._manager._dispatch_record
@@ -5,7 +5,7 @@ from dataclasses import dataclass
5
5
  from typing import Any
6
6
 
7
7
  from statewire import Statewire, command
8
- from statewire_helpers import apply_ops, post_command, statewire_client, stream_of
8
+ from statewire_helpers import apply_ops, attach, post_command, statewire_client, stream_of
9
9
 
10
10
  from harness_sdk import RunManager
11
11
 
@@ -32,6 +32,7 @@ class Script:
32
32
  def __init__(self) -> None:
33
33
  self.calls: asyncio.Queue[Call] = asyncio.Queue()
34
34
  self.thread: dict[str, dict[str, Any]] = {}
35
+ self.leaf: str | None = None
35
36
 
36
37
  async def run(self, ctx: RunManager.RunContext) -> Any:
37
38
  call = Call(ctx, asyncio.get_running_loop().create_future())
@@ -43,6 +44,9 @@ class Script:
43
44
  return {"isLeaf": not self.thread}
44
45
  return self.thread.get(message_id)
45
46
 
47
+ async def get_leaf_message_id(self) -> str | None:
48
+ return self.leaf
49
+
46
50
  async def next_call(self, timeout: float = 5) -> Call:
47
51
  return await asyncio.wait_for(self.calls.get(), timeout)
48
52
 
@@ -55,7 +59,14 @@ def _meta(params: Any) -> Any:
55
59
  return params.pop("meta", None) if isinstance(params, dict) else None
56
60
 
57
61
 
58
- def make_host(script: Script, capabilities=(), initial_runs=None, max_queued=50):
62
+ def make_host(
63
+ script: Script,
64
+ capabilities=(),
65
+ initial_runs=None,
66
+ max_queued=50,
67
+ prepare_message=None,
68
+ prepare_input=None,
69
+ ):
59
70
  class Host(Statewire):
60
71
  live: "Host | None" = None
61
72
 
@@ -67,11 +78,13 @@ def make_host(script: Script, capabilities=(), initial_runs=None, max_queued=50)
67
78
  self.runs = RunManager(
68
79
  state=self.state,
69
80
  run=script.run,
70
- get_message_meta=script.get_message_meta,
81
+ thread=script,
71
82
  create_task=self.create_task,
72
83
  schedule=self.schedule,
73
84
  capabilities=capabilities,
74
85
  max_queued=max_queued,
86
+ prepare_message=prepare_message,
87
+ prepare_input=prepare_input,
75
88
  )
76
89
  yield
77
90
 
@@ -114,20 +127,47 @@ class RunDriver:
114
127
  """One SSE-attached statewire client: sends commands, mirrors the state
115
128
  replica from ops, and returns each command's settled response."""
116
129
 
117
- def __init__(self, client, stream) -> None:
130
+ def __init__(self, client, stream, app=None) -> None:
118
131
  self._client = client
119
132
  self._stream = stream
133
+ self._app = app
134
+ self._side_seqs: dict[str, int] = {}
120
135
  self._seq = 0
121
136
  self._acked = 0
122
137
  self._res: dict[int, list[dict[str, Any]]] = {}
138
+ self._posts: list["asyncio.Task[None]"] = []
123
139
  self.replica: dict[str, Any] = {}
124
140
  self.envelopes: list[dict[str, Any]] = []
125
141
 
142
+ def _spawn(self, body: dict[str, Any], seq: int) -> None:
143
+ async def _post() -> None:
144
+ response = await post_command(self._client, body, seq=seq)
145
+ assert response.status_code == 200, response.text
146
+
147
+ self._posts.append(asyncio.create_task(_post()))
148
+
149
+ def post(self, method: str, params: Any = None) -> int:
150
+ """Background POST for a command whose receipt is held until its effect
151
+ publishes (a dispatching send or rewind); settle via ``res(seq)``."""
152
+ self._seq += 1
153
+ body = {"method": method, "params": [] if params is None else [params]}
154
+ self._spawn(body, self._seq)
155
+ return self._seq
156
+
157
+ async def side(self, method: str, params: Any = None, *, client_id: str = "c2"):
158
+ """A second client's command: its own seq stream stays admissible while
159
+ this client's receipt is deferred; its verdicts ride the other stream."""
160
+ await attach(self._app, client_id=client_id)
161
+ seq = self._side_seqs[client_id] = self._side_seqs.get(client_id, 0) + 1
162
+ body = {"method": method, "params": [] if params is None else [params]}
163
+ response = await post_command(self._client, body, seq=seq, client_id=client_id)
164
+ assert response.status_code == 200, response.text
165
+
126
166
  async def pump(self, timeout: float = 5) -> None:
127
167
  env = await self._stream.next_event(timeout=timeout)
128
168
  assert env is not None, "stream ended unexpectedly"
129
169
  self.envelopes.append(env)
130
- apply_ops(self.replica, env)
170
+ apply_ops(self.replica, copy.deepcopy(env))
131
171
  if "ack" in env:
132
172
  self._acked = max(self._acked, env["ack"])
133
173
  for rsp in env.get("res", []):
@@ -143,9 +183,9 @@ class RunDriver:
143
183
  assert response.status_code == 200, response.text
144
184
  return await self.res(seq, terminal=terminal)
145
185
 
146
- async def batch(self, commands: list[tuple[str, Any]]) -> int:
147
- """POST one batch; returns the first member's seq (members are
148
- consecutive). Responses via ``res(first + i)``."""
186
+ def batch(self, commands: list[tuple[str, Any]]) -> int:
187
+ """Background POST of one batch; returns the first member's seq
188
+ (members are consecutive). Responses via ``res(first + i)``."""
149
189
  first = self._seq + 1
150
190
  self._seq += len(commands)
151
191
  body = {
@@ -154,8 +194,7 @@ class RunDriver:
154
194
  for method, params in commands
155
195
  ]
156
196
  }
157
- response = await post_command(self._client, body, seq=first)
158
- assert response.status_code == 200, response.text
197
+ self._spawn(body, first)
159
198
  return first
160
199
 
161
200
  async def context(self, value: dict[str, Any]) -> dict[str, Any]:
@@ -218,19 +257,38 @@ class RunDriver:
218
257
 
219
258
 
220
259
  @asynccontextmanager
221
- async def run_host(script: Script, capabilities=(), initial_runs=None, max_queued=50):
260
+ async def run_host(
261
+ script: Script,
262
+ capabilities=(),
263
+ initial_runs=None,
264
+ max_queued=50,
265
+ prepare_message=None,
266
+ prepare_input=None,
267
+ ):
222
268
  host_cls = make_host(
223
269
  script,
224
270
  capabilities=capabilities,
225
271
  initial_runs=initial_runs,
226
272
  max_queued=max_queued,
273
+ prepare_message=prepare_message,
274
+ prepare_input=prepare_input,
227
275
  )
228
276
  async with statewire_client(host_cls) as (app, client):
229
277
  async with stream_of(app) as stream:
230
278
  snapshot = await stream.next_event()
231
- driver = RunDriver(client, stream)
232
- apply_ops(driver.replica, snapshot)
233
- yield driver, host_cls
279
+ driver = RunDriver(client, stream, app)
280
+ driver.envelopes.append(snapshot)
281
+ apply_ops(driver.replica, copy.deepcopy(snapshot))
282
+ try:
283
+ yield driver, host_cls
284
+ finally:
285
+ for task in driver._posts:
286
+ task.cancel()
287
+ for result in await asyncio.gather(
288
+ *driver._posts, return_exceptions=True
289
+ ):
290
+ if isinstance(result, AssertionError):
291
+ raise result
234
292
 
235
293
 
236
294
  def msg(id: str, text: str | None = None) -> dict[str, Any]: