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.
- {harness_sdk_python-0.5.0 → harness_sdk_python-0.7.0}/.gitignore +1 -0
- {harness_sdk_python-0.5.0 → harness_sdk_python-0.7.0}/PKG-INFO +2 -2
- {harness_sdk_python-0.5.0 → harness_sdk_python-0.7.0}/README.md +1 -1
- {harness_sdk_python-0.5.0 → harness_sdk_python-0.7.0}/pyproject.toml +1 -1
- harness_sdk_python-0.7.0/src/harness_sdk/__init__.py +4 -0
- {harness_sdk_python-0.5.0 → harness_sdk_python-0.7.0}/src/harness_sdk/fenced_postgres.py +58 -18
- harness_sdk_python-0.7.0/src/harness_sdk/linear_thread.py +52 -0
- {harness_sdk_python-0.5.0 → harness_sdk_python-0.7.0}/src/harness_sdk/run_manager.py +79 -34
- {harness_sdk_python-0.5.0 → harness_sdk_python-0.7.0}/tests/run_helpers.py +72 -14
- harness_sdk_python-0.7.0/tests/test_ack_visibility.py +83 -0
- {harness_sdk_python-0.5.0 → harness_sdk_python-0.7.0}/tests/test_batches.py +12 -11
- {harness_sdk_python-0.5.0 → harness_sdk_python-0.7.0}/tests/test_branch_anchor.py +12 -17
- {harness_sdk_python-0.5.0 → harness_sdk_python-0.7.0}/tests/test_edit_dispatched.py +14 -23
- {harness_sdk_python-0.5.0 → harness_sdk_python-0.7.0}/tests/test_edit_reload.py +19 -28
- {harness_sdk_python-0.5.0 → harness_sdk_python-0.7.0}/tests/test_enqueue.py +28 -30
- {harness_sdk_python-0.5.0 → harness_sdk_python-0.7.0}/tests/test_facade.py +1 -1
- {harness_sdk_python-0.5.0 → harness_sdk_python-0.7.0}/tests/test_fenced_postgres.py +64 -0
- {harness_sdk_python-0.5.0 → harness_sdk_python-0.7.0}/tests/test_input_required.py +16 -18
- harness_sdk_python-0.7.0/tests/test_linear_thread.py +122 -0
- {harness_sdk_python-0.5.0 → harness_sdk_python-0.7.0}/tests/test_meta.py +16 -22
- {harness_sdk_python-0.5.0 → harness_sdk_python-0.7.0}/tests/test_outcomes.py +15 -14
- {harness_sdk_python-0.5.0 → harness_sdk_python-0.7.0}/tests/test_placement.py +2 -2
- harness_sdk_python-0.7.0/tests/test_prepare_hooks.py +164 -0
- {harness_sdk_python-0.5.0 → harness_sdk_python-0.7.0}/tests/test_rewind_during_run.py +14 -28
- {harness_sdk_python-0.5.0 → harness_sdk_python-0.7.0}/tests/test_run_leaf.py +53 -48
- {harness_sdk_python-0.5.0 → harness_sdk_python-0.7.0}/tests/test_settle.py +36 -60
- {harness_sdk_python-0.5.0 → harness_sdk_python-0.7.0}/tests/test_steer.py +16 -12
- {harness_sdk_python-0.5.0 → harness_sdk_python-0.7.0}/tests/test_stop_continue.py +28 -24
- harness_sdk_python-0.5.0/src/harness_sdk/__init__.py +0 -4
- harness_sdk_python-0.5.0/src/harness_sdk/linear_thread.py +0 -31
- harness_sdk_python-0.5.0/tests/test_linear_thread.py +0 -87
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: harness-sdk-python
|
|
3
|
-
Version: 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
|
-
|
|
26
|
+
thread=self._thread,
|
|
27
27
|
create_task=self.create_task,
|
|
28
28
|
capabilities=("rewind",),
|
|
29
29
|
)
|
|
@@ -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
|
-
|
|
153
|
-
|
|
154
|
-
|
|
155
|
-
|
|
156
|
-
|
|
157
|
-
|
|
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(
|
|
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
|
-
|
|
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.
|
|
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
|
-
|
|
316
|
-
|
|
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.
|
|
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.
|
|
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.
|
|
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.
|
|
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.
|
|
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.
|
|
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.
|
|
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(
|
|
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
|
-
|
|
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
|
-
|
|
147
|
-
"""POST one batch; returns the first member's seq
|
|
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
|
-
|
|
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(
|
|
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
|
-
|
|
233
|
-
|
|
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]:
|