statewire 0.3.2__tar.gz → 0.4.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.
- {statewire-0.3.2 → statewire-0.4.0}/PKG-INFO +1 -1
- {statewire-0.3.2 → statewire-0.4.0}/pyproject.toml +1 -1
- {statewire-0.3.2 → statewire-0.4.0}/src/statewire/assistant_transport.py +148 -66
- {statewire-0.3.2 → statewire-0.4.0}/tests/test_assistant_transport.py +211 -11
- {statewire-0.3.2 → statewire-0.4.0}/tests/test_assistant_transport_client.py +2 -2
- {statewire-0.3.2 → statewire-0.4.0}/.gitignore +0 -0
- {statewire-0.3.2 → statewire-0.4.0}/README.md +0 -0
- {statewire-0.3.2 → statewire-0.4.0}/examples/__init__.py +0 -0
- {statewire-0.3.2 → statewire-0.4.0}/examples/demo_app.py +0 -0
- {statewire-0.3.2 → statewire-0.4.0}/src/statewire/__init__.py +0 -0
- {statewire-0.3.2 → statewire-0.4.0}/src/statewire/api.py +0 -0
- {statewire-0.3.2 → statewire-0.4.0}/src/statewire/assistant_transport_client.py +0 -0
- {statewire-0.3.2 → statewire-0.4.0}/src/statewire/client.py +0 -0
- {statewire-0.3.2 → statewire-0.4.0}/src/statewire/langgraph.py +0 -0
- {statewire-0.3.2 → statewire-0.4.0}/src/statewire/ops.py +0 -0
- {statewire-0.3.2 → statewire-0.4.0}/src/statewire/state.py +0 -0
- {statewire-0.3.2 → statewire-0.4.0}/tests/client_helpers.py +0 -0
- {statewire-0.3.2 → statewire-0.4.0}/tests/statewire_helpers.py +0 -0
- {statewire-0.3.2 → statewire-0.4.0}/tests/test_assistant_transport_facade.py +0 -0
- {statewire-0.3.2 → statewire-0.4.0}/tests/test_authorize.py +0 -0
- {statewire-0.3.2 → statewire-0.4.0}/tests/test_client.py +0 -0
- {statewire-0.3.2 → statewire-0.4.0}/tests/test_client_ws.py +0 -0
- {statewire-0.3.2 → statewire-0.4.0}/tests/test_commands.py +0 -0
- {statewire-0.3.2 → statewire-0.4.0}/tests/test_context.py +0 -0
- {statewire-0.3.2 → statewire-0.4.0}/tests/test_langgraph.py +0 -0
- {statewire-0.3.2 → statewire-0.4.0}/tests/test_lifespan.py +0 -0
- {statewire-0.3.2 → statewire-0.4.0}/tests/test_meta.py +0 -0
- {statewire-0.3.2 → statewire-0.4.0}/tests/test_state_proxy.py +0 -0
- {statewire-0.3.2 → statewire-0.4.0}/tests/test_statewire_hostable.py +0 -0
- {statewire-0.3.2 → statewire-0.4.0}/tests/test_stream.py +0 -0
- {statewire-0.3.2 → statewire-0.4.0}/tests/test_writer_lease.py +0 -0
- {statewire-0.3.2 → statewire-0.4.0}/tests/test_ws.py +0 -0
|
@@ -17,32 +17,38 @@ the legacy commands, which are normalized to the canonical ``run/*`` dialect
|
|
|
17
17
|
(a host-registered handler of the same legacy name always wins and receives
|
|
18
18
|
the command whole):
|
|
19
19
|
|
|
20
|
-
- ``add-message`` becomes ``run/steer`` with ``{"message": <wire message
|
|
21
|
-
|
|
20
|
+
- ``add-message`` becomes ``run/steer`` with ``{"message": <wire message>,
|
|
21
|
+
"anchorMessageId": <body parentId>}`` (body ``parentId`` null or absent →
|
|
22
|
+
``null``): a fresh ``legacy_`` id is synthesized, image parts become
|
|
22
23
|
``{"type": "file", "mediaType": "image/*", "url"}`` parts, and command-level
|
|
23
|
-
anchors are dropped.
|
|
24
|
-
``
|
|
25
|
-
``
|
|
24
|
+
anchors are dropped. A command-level non-null ``sourceId`` makes it
|
|
25
|
+
``run/edit`` with that ``sourceId`` directly; otherwise, when the body
|
|
26
|
+
carries a top-level ``parentId`` and ``get_message_child_id`` resolves it to
|
|
27
|
+
a child, the command becomes ``run/edit`` with ``{"sourceId": <child>,
|
|
28
|
+
"message": <wire message>}``.
|
|
26
29
|
- ``add-tool-result`` becomes ``run/input`` with ``{"requestId": <resolved>,
|
|
27
30
|
"response": {"output": <result>, "isError"?, "artifact"?}}``; the requestId
|
|
28
31
|
comes from ``get_input_request_id(toolCallId)``, whose default peeks the
|
|
29
32
|
unanswered tool-call request in ``state["inputRequests"]`` (None rejects),
|
|
30
33
|
and a ``modelContent`` field rejects (unsupported in ``run/input``).
|
|
31
34
|
- A body with a ``parentId`` and no commands is the legacy reload: it becomes
|
|
32
|
-
``run/reload`` with ``{"sourceId": <child
|
|
33
|
-
(None rejects).
|
|
35
|
+
``run/reload`` with ``{"sourceId": <child>, "messageId": <fresh legacy_
|
|
36
|
+
id>}`` via ``get_message_child_id`` (None rejects).
|
|
34
37
|
|
|
35
38
|
A host without the target ``run/*`` handler rejects the batch via the
|
|
36
|
-
ordinary unknown-command path. The response streams
|
|
37
|
-
|
|
38
|
-
|
|
39
|
-
|
|
40
|
-
|
|
41
|
-
|
|
39
|
+
ordinary unknown-command path. The response streams state changes as legacy
|
|
40
|
+
frames until the instance is idle — every submitted command settled and no
|
|
41
|
+
live instance task (``create_task`` minus ``unref``) — then EOF. A run error
|
|
42
|
+
recorded into state (``get_run_error``; default ``state["error"]``) becomes
|
|
43
|
+
an error frame before EOF, its fields beyond ``message`` carried as the
|
|
44
|
+
frame's structured payload where the format allows; a rejection or crash
|
|
45
|
+
becomes an error frame followed by EOF. Client disconnect before EOF invokes
|
|
46
|
+
``on_assistant_transport_disconnect`` (a plain overridable hook, not a
|
|
47
|
+
command; default no-op); the run itself keeps executing server-side.
|
|
42
48
|
|
|
43
49
|
``POST /assistant-transport/api/resume`` (same body shape, ``commands`` must
|
|
44
50
|
be absent or empty) reattaches to the run: a full-state ``set`` at the root,
|
|
45
|
-
then the live tail until the
|
|
51
|
+
then the live tail until the instance is idle. If no run ever started, the
|
|
46
52
|
response is 200 with an empty body and ``X-Stream-Status: not_found``. If the
|
|
47
53
|
run already completed, the replay carries ``X-Stream-Status: completed``. A
|
|
48
54
|
resume disconnect never invokes the cancel hook.
|
|
@@ -56,7 +62,7 @@ existed at request time.
|
|
|
56
62
|
``POST /assistant-transport/api/status`` responds 200 always:
|
|
57
63
|
``{"isRunning": true, "status": "running"}`` while a run is active,
|
|
58
64
|
``{"isRunning": false, "status": "completed", "completedAt": <ms epoch>}``
|
|
59
|
-
|
|
65
|
+
once idle, and ``{"isRunning": false, "status": "not_found",
|
|
60
66
|
"message": ...}`` when no run ever started on this instance.
|
|
61
67
|
|
|
62
68
|
Two stream formats, selected by ``assistant_transport_protocol``:
|
|
@@ -74,7 +80,7 @@ import contextlib
|
|
|
74
80
|
import json
|
|
75
81
|
import secrets
|
|
76
82
|
import time
|
|
77
|
-
from typing import Any, AsyncIterator
|
|
83
|
+
from typing import Any, AsyncIterator, Coroutine
|
|
78
84
|
from uuid import uuid4
|
|
79
85
|
|
|
80
86
|
import httpx
|
|
@@ -102,7 +108,7 @@ class _DataStreamFormat:
|
|
|
102
108
|
return f"aui-state:[{','.join(op_frames)}]\n"
|
|
103
109
|
|
|
104
110
|
@staticmethod
|
|
105
|
-
def error(message: str) -> str:
|
|
111
|
+
def error(message: str, payload: Any = None) -> str:
|
|
106
112
|
return f"3:{json.dumps(message)}\n"
|
|
107
113
|
|
|
108
114
|
|
|
@@ -119,9 +125,11 @@ class _AssistantTransportFormat:
|
|
|
119
125
|
)
|
|
120
126
|
|
|
121
127
|
@staticmethod
|
|
122
|
-
def error(message: str) -> str:
|
|
123
|
-
|
|
124
|
-
|
|
128
|
+
def error(message: str, payload: Any = None) -> str:
|
|
129
|
+
chunk: dict[str, Any] = {"type": "error", "error": message}
|
|
130
|
+
if payload is not None:
|
|
131
|
+
chunk["payload"] = payload
|
|
132
|
+
return f"data: {json.dumps(chunk)}\n\n"
|
|
125
133
|
|
|
126
134
|
|
|
127
135
|
_FORMATS = {
|
|
@@ -183,7 +191,9 @@ _AT_CONSUMED_FIELDS = {"commands", "threadId", "state"}
|
|
|
183
191
|
|
|
184
192
|
|
|
185
193
|
class _RunError(Exception):
|
|
186
|
-
|
|
194
|
+
def __init__(self, message: str, *, payload: Any = None) -> None:
|
|
195
|
+
super().__init__(message)
|
|
196
|
+
self.payload = payload
|
|
187
197
|
|
|
188
198
|
|
|
189
199
|
def _legacy_wire_message(message: Any) -> dict[str, Any]:
|
|
@@ -215,6 +225,7 @@ class _Run:
|
|
|
215
225
|
def __init__(self, replica: Any) -> None:
|
|
216
226
|
self.replica = replica
|
|
217
227
|
self.error: str | None = None
|
|
228
|
+
self.error_payload: Any = None
|
|
218
229
|
self.done = False
|
|
219
230
|
self.completed_at: int | None = None
|
|
220
231
|
self.attachments: list["asyncio.Queue[_Item]"] = []
|
|
@@ -227,7 +238,7 @@ class _Run:
|
|
|
227
238
|
self.attachments.remove(aq)
|
|
228
239
|
while not aq.empty():
|
|
229
240
|
aq.get_nowait()
|
|
230
|
-
aq.put_nowait(("error", "state stream overflowed"))
|
|
241
|
+
aq.put_nowait(("error", ("state stream overflowed", None)))
|
|
231
242
|
aq.put_nowait(("end", None))
|
|
232
243
|
|
|
233
244
|
|
|
@@ -236,6 +247,34 @@ class AssistantTransport(Statewire):
|
|
|
236
247
|
|
|
237
248
|
_at_run: _Run | None = None
|
|
238
249
|
|
|
250
|
+
def __init__(self, ctx: Any) -> None:
|
|
251
|
+
super().__init__(ctx)
|
|
252
|
+
self._at_tasks: set["asyncio.Task[Any]"] = set()
|
|
253
|
+
|
|
254
|
+
def create_task(self, coro: Coroutine[Any, Any, Any]) -> "asyncio.Task[Any]":
|
|
255
|
+
task = super().create_task(coro)
|
|
256
|
+
self._at_tasks.add(task)
|
|
257
|
+
task.add_done_callback(self._at_tasks.discard)
|
|
258
|
+
return task
|
|
259
|
+
|
|
260
|
+
def unref(self, task: "asyncio.Task[Any]") -> "asyncio.Task[Any]":
|
|
261
|
+
self._at_tasks.discard(task)
|
|
262
|
+
return super().unref(task)
|
|
263
|
+
|
|
264
|
+
async def get_run_error(self) -> dict[str, Any] | None:
|
|
265
|
+
"""The run error to surface as the end-of-stream error frame.
|
|
266
|
+
|
|
267
|
+
The default peeks the RunManager root-state mount: a non-null
|
|
268
|
+
``state["error"]`` object with a string ``message`` is the recorded
|
|
269
|
+
run error; its fields beyond ``message`` travel as the frame's
|
|
270
|
+
structured payload where the format carries one. None ends the
|
|
271
|
+
stream cleanly."""
|
|
272
|
+
state = plain(self.state)
|
|
273
|
+
error = state.get("error") if isinstance(state, dict) else None
|
|
274
|
+
if isinstance(error, dict) and isinstance(error.get("message"), str):
|
|
275
|
+
return error
|
|
276
|
+
return None
|
|
277
|
+
|
|
239
278
|
async def on_assistant_transport_disconnect(self) -> None:
|
|
240
279
|
"""The legacy client went away or asked to stop; treat as a cancel signal."""
|
|
241
280
|
|
|
@@ -308,11 +347,12 @@ class AssistantTransport(Statewire):
|
|
|
308
347
|
parent_id = None
|
|
309
348
|
run = _Run(json.loads(snapshot)["ops"][0]["value"])
|
|
310
349
|
self._at_run = run
|
|
311
|
-
self.create_task(
|
|
350
|
+
watcher = self.create_task(
|
|
312
351
|
self._assistant_transport_run(
|
|
313
352
|
run, q, client_id, lease, commands, parent_id
|
|
314
353
|
)
|
|
315
354
|
)
|
|
355
|
+
self._at_tasks.discard(watcher)
|
|
316
356
|
return StreamingResponse(
|
|
317
357
|
self._assistant_transport_attach(fmt, run, initial=True),
|
|
318
358
|
media_type="text/event-stream",
|
|
@@ -415,17 +455,20 @@ class AssistantTransport(Statewire):
|
|
|
415
455
|
registered = type(self)._statewire_commands
|
|
416
456
|
if kind == "add-message" and "add-message" not in registered:
|
|
417
457
|
message = _legacy_wire_message(command.get("message"))
|
|
418
|
-
source_id = (
|
|
419
|
-
|
|
420
|
-
|
|
421
|
-
|
|
422
|
-
|
|
458
|
+
source_id = command.get("sourceId")
|
|
459
|
+
if source_id is not None and not isinstance(source_id, str):
|
|
460
|
+
raise _RunError("add-message: sourceId must be a string or null")
|
|
461
|
+
if source_id is None and parent_id is not None:
|
|
462
|
+
source_id = await self.get_message_child_id(parent_id)
|
|
423
463
|
if source_id is not None:
|
|
424
464
|
return {
|
|
425
465
|
"method": "run/edit",
|
|
426
466
|
"params": [{"sourceId": source_id, "message": message}],
|
|
427
467
|
}
|
|
428
|
-
return {
|
|
468
|
+
return {
|
|
469
|
+
"method": "run/steer",
|
|
470
|
+
"params": [{"message": message, "anchorMessageId": parent_id}],
|
|
471
|
+
}
|
|
429
472
|
if kind == "add-tool-result" and "add-tool-result" not in registered:
|
|
430
473
|
return {
|
|
431
474
|
"method": "run/input",
|
|
@@ -455,7 +498,12 @@ class AssistantTransport(Statewire):
|
|
|
455
498
|
source_id = await self.get_message_child_id(parent_id)
|
|
456
499
|
if source_id is None:
|
|
457
500
|
raise _RunError(f"reload: no child message for parentId {parent_id!r}")
|
|
458
|
-
return {
|
|
501
|
+
return {
|
|
502
|
+
"method": "run/reload",
|
|
503
|
+
"params": [
|
|
504
|
+
{"sourceId": source_id, "messageId": f"legacy_{uuid4().hex}"}
|
|
505
|
+
],
|
|
506
|
+
}
|
|
459
507
|
|
|
460
508
|
async def _assistant_transport_dispatch(self, method: str, identity: Any) -> None:
|
|
461
509
|
client_id = "at-" + secrets.token_urlsafe(9)
|
|
@@ -486,48 +534,82 @@ class AssistantTransport(Statewire):
|
|
|
486
534
|
total = len(commands) or (1 if parent_id is not None else 0)
|
|
487
535
|
acked = 0
|
|
488
536
|
open_pending: set[int] = set()
|
|
537
|
+
|
|
538
|
+
def idle() -> bool:
|
|
539
|
+
return (
|
|
540
|
+
submit.done()
|
|
541
|
+
and acked >= total
|
|
542
|
+
and not open_pending
|
|
543
|
+
and not self._at_tasks
|
|
544
|
+
)
|
|
545
|
+
|
|
546
|
+
def consume(item: Any) -> bool:
|
|
547
|
+
nonlocal acked
|
|
548
|
+
if item is _LAGGARD:
|
|
549
|
+
raise _RunError("state stream overflowed")
|
|
550
|
+
data, finish = item
|
|
551
|
+
envelope = json.loads(data)
|
|
552
|
+
op_frames: list[str] = []
|
|
553
|
+
for op in envelope.get("ops", []):
|
|
554
|
+
run.replica, frame = _translate_op(run.replica, op)
|
|
555
|
+
if frame is not None:
|
|
556
|
+
op_frames.append(frame)
|
|
557
|
+
if op_frames:
|
|
558
|
+
run.push(("state", op_frames))
|
|
559
|
+
for res in envelope.get("res", []):
|
|
560
|
+
if res["type"] == "pending":
|
|
561
|
+
open_pending.add(res["seq"])
|
|
562
|
+
continue
|
|
563
|
+
open_pending.discard(res["seq"])
|
|
564
|
+
if res["type"] in ("rejected", "crashed"):
|
|
565
|
+
raise _RunError(
|
|
566
|
+
res.get("message") or res["type"],
|
|
567
|
+
payload=res.get("payload"),
|
|
568
|
+
)
|
|
569
|
+
if "ack" in envelope:
|
|
570
|
+
acked = envelope["ack"]
|
|
571
|
+
if finish and not idle():
|
|
572
|
+
fin = envelope.get("fin", {})
|
|
573
|
+
raise _RunError(
|
|
574
|
+
fin.get("message") or fin.get("reason", "stream ended")
|
|
575
|
+
)
|
|
576
|
+
return finish
|
|
577
|
+
|
|
489
578
|
try:
|
|
490
|
-
|
|
491
|
-
|
|
579
|
+
ended = False
|
|
580
|
+
while not ended:
|
|
581
|
+
if idle():
|
|
582
|
+
self.drain()
|
|
583
|
+
self.flush()
|
|
584
|
+
if q_get.done():
|
|
585
|
+
ended = consume(q_get.result()) or ended
|
|
586
|
+
q_get = asyncio.create_task(q.get())
|
|
587
|
+
continue
|
|
588
|
+
while not q.empty():
|
|
589
|
+
ended = consume(q.get_nowait()) or ended
|
|
590
|
+
if idle():
|
|
591
|
+
break
|
|
592
|
+
continue
|
|
593
|
+
waiters = {q_get, *self._at_tasks}
|
|
594
|
+
if not submit.done():
|
|
595
|
+
waiters.add(submit)
|
|
492
596
|
done, _ = await asyncio.wait(
|
|
493
597
|
waiters, return_when=asyncio.FIRST_COMPLETED
|
|
494
598
|
)
|
|
495
599
|
if submit in done and submit.exception() is not None:
|
|
496
600
|
raise _RunError(str(submit.exception()))
|
|
497
|
-
if q_get
|
|
498
|
-
|
|
499
|
-
|
|
500
|
-
|
|
501
|
-
|
|
502
|
-
|
|
503
|
-
|
|
504
|
-
|
|
505
|
-
for op in envelope.get("ops", []):
|
|
506
|
-
run.replica, frame = _translate_op(run.replica, op)
|
|
507
|
-
if frame is not None:
|
|
508
|
-
op_frames.append(frame)
|
|
509
|
-
if op_frames:
|
|
510
|
-
run.push(("state", op_frames))
|
|
511
|
-
for res in envelope.get("res", []):
|
|
512
|
-
if res["type"] == "pending":
|
|
513
|
-
open_pending.add(res["seq"])
|
|
514
|
-
continue
|
|
515
|
-
open_pending.discard(res["seq"])
|
|
516
|
-
if res["type"] in ("rejected", "crashed"):
|
|
517
|
-
raise _RunError(res.get("message") or res["type"])
|
|
518
|
-
if "ack" in envelope:
|
|
519
|
-
acked = envelope["ack"]
|
|
520
|
-
if finish:
|
|
521
|
-
if submit.done() and acked >= total and not open_pending:
|
|
522
|
-
break
|
|
523
|
-
fin = envelope.get("fin", {})
|
|
524
|
-
raise _RunError(
|
|
525
|
-
fin.get("message") or fin.get("reason", "stream ended")
|
|
526
|
-
)
|
|
527
|
-
q_get = asyncio.create_task(q.get())
|
|
601
|
+
if q_get in done:
|
|
602
|
+
ended = consume(q_get.result())
|
|
603
|
+
if not ended:
|
|
604
|
+
q_get = asyncio.create_task(q.get())
|
|
605
|
+
error = await self.get_run_error()
|
|
606
|
+
if error is not None:
|
|
607
|
+
payload = {k: v for k, v in error.items() if k != "message"}
|
|
608
|
+
raise _RunError(error["message"], payload=payload or None)
|
|
528
609
|
except _RunError as exc:
|
|
529
610
|
run.error = str(exc)
|
|
530
|
-
run.
|
|
611
|
+
run.error_payload = exc.payload
|
|
612
|
+
run.push(("error", (run.error, run.error_payload)))
|
|
531
613
|
finally:
|
|
532
614
|
q_get.cancel()
|
|
533
615
|
submit.cancel()
|
|
@@ -559,11 +641,11 @@ class AssistantTransport(Statewire):
|
|
|
559
641
|
if kind == "state":
|
|
560
642
|
yield fmt.state(value)
|
|
561
643
|
elif kind == "error":
|
|
562
|
-
yield fmt.error(value)
|
|
644
|
+
yield fmt.error(*value)
|
|
563
645
|
else:
|
|
564
646
|
break
|
|
565
647
|
elif initial and run.error is not None:
|
|
566
|
-
yield fmt.error(run.error)
|
|
648
|
+
yield fmt.error(run.error, run.error_payload)
|
|
567
649
|
if fmt.end:
|
|
568
650
|
yield fmt.end
|
|
569
651
|
completed = True
|
|
@@ -2,13 +2,16 @@
|
|
|
2
2
|
|
|
3
3
|
The legacy client POSTs ``{"commands", "state", "threadId", ...}`` and reads
|
|
4
4
|
back a data-stream (``<prefix>:<json>`` lines) or assistant-transport SSE
|
|
5
|
-
response carrying ``set``/``append-text`` state ops until the
|
|
5
|
+
response carrying ``set``/``append-text`` state ops until the instance is
|
|
6
|
+
idle.
|
|
6
7
|
"""
|
|
7
8
|
|
|
8
9
|
import asyncio
|
|
9
10
|
import contextlib
|
|
10
11
|
import json
|
|
11
12
|
|
|
13
|
+
import pytest
|
|
14
|
+
|
|
12
15
|
from statewire_helpers import statewire_client
|
|
13
16
|
|
|
14
17
|
from statewire import AssistantTransport, StatewireReject, command
|
|
@@ -24,7 +27,7 @@ class Harness(AssistantTransport):
|
|
|
24
27
|
self.state = {"messages": [], "text": "", "meta": {"title": None}}
|
|
25
28
|
yield
|
|
26
29
|
|
|
27
|
-
@command("
|
|
30
|
+
@command("demo/echo")
|
|
28
31
|
async def add_message(self, cmd):
|
|
29
32
|
self.seen.append(cmd)
|
|
30
33
|
self.state["messages"].append(cmd["message"])
|
|
@@ -110,14 +113,14 @@ async def post_run(client, body, thread="t1"):
|
|
|
110
113
|
return await client.post(f"/threads/{thread}/assistant-transport/api/chat", json=body)
|
|
111
114
|
|
|
112
115
|
|
|
113
|
-
async def
|
|
116
|
+
async def test_data_stream_response_streams_snapshot_then_ops_until_idle():
|
|
114
117
|
async with statewire_client(Harness) as (app, client):
|
|
115
118
|
response = await post_run(
|
|
116
119
|
client,
|
|
117
120
|
{
|
|
118
121
|
"commands": [
|
|
119
122
|
{
|
|
120
|
-
"type": "
|
|
123
|
+
"type": "demo/echo",
|
|
121
124
|
"message": {
|
|
122
125
|
"role": "user",
|
|
123
126
|
"parts": [{"type": "text", "text": "hi"}],
|
|
@@ -226,7 +229,7 @@ async def test_splice_and_remove_fall_back_to_parent_set():
|
|
|
226
229
|
{
|
|
227
230
|
"commands": [
|
|
228
231
|
{
|
|
229
|
-
"type": "
|
|
232
|
+
"type": "demo/echo",
|
|
230
233
|
"message": {
|
|
231
234
|
"role": "user",
|
|
232
235
|
"parts": [{"type": "text", "text": "hi"}],
|
|
@@ -263,7 +266,7 @@ async def test_assistant_transport_sse_mode():
|
|
|
263
266
|
{
|
|
264
267
|
"commands": [
|
|
265
268
|
{
|
|
266
|
-
"type": "
|
|
269
|
+
"type": "demo/echo",
|
|
267
270
|
"message": {
|
|
268
271
|
"role": "user",
|
|
269
272
|
"parts": [{"type": "text", "text": "hi"}],
|
|
@@ -347,6 +350,14 @@ class _PostStream:
|
|
|
347
350
|
async def next_chunk(self, timeout: float = 5) -> bytes:
|
|
348
351
|
return await asyncio.wait_for(self._chunks.get(), timeout)
|
|
349
352
|
|
|
353
|
+
async def until_closed(self, timeout: float = 5) -> bytes:
|
|
354
|
+
assert self._task is not None
|
|
355
|
+
await asyncio.wait_for(asyncio.shield(self._task), timeout)
|
|
356
|
+
tail = b""
|
|
357
|
+
while not self._chunks.empty():
|
|
358
|
+
tail += self._chunks.get_nowait()
|
|
359
|
+
return tail
|
|
360
|
+
|
|
350
361
|
async def __aexit__(self, *exc):
|
|
351
362
|
self._disconnected.set()
|
|
352
363
|
assert self._task is not None
|
|
@@ -417,7 +428,7 @@ async def test_resume_after_completed_run_replays_with_completed_header():
|
|
|
417
428
|
{
|
|
418
429
|
"commands": [
|
|
419
430
|
{
|
|
420
|
-
"type": "
|
|
431
|
+
"type": "demo/echo",
|
|
421
432
|
"message": {
|
|
422
433
|
"role": "user",
|
|
423
434
|
"parts": [{"type": "text", "text": "hi"}],
|
|
@@ -442,7 +453,7 @@ async def test_resume_rejects_commands_and_malformed_bodies():
|
|
|
442
453
|
async with statewire_client(Harness) as (app, client):
|
|
443
454
|
response = await client.post(
|
|
444
455
|
"/threads/t1/assistant-transport/api/resume",
|
|
445
|
-
json={"commands": [{"type": "
|
|
456
|
+
json={"commands": [{"type": "demo/echo"}]},
|
|
446
457
|
)
|
|
447
458
|
assert response.status_code == 400
|
|
448
459
|
assert (
|
|
@@ -597,6 +608,138 @@ async def test_status_reports_not_found_running_and_completed():
|
|
|
597
608
|
assert isinstance(body["completedAt"], int)
|
|
598
609
|
|
|
599
610
|
|
|
611
|
+
class BackgroundHarness(AssistantTransport):
|
|
612
|
+
"""Host whose run/steer settles at enqueue; the run is an instance task."""
|
|
613
|
+
|
|
614
|
+
def __init__(self, ctx):
|
|
615
|
+
super().__init__(ctx)
|
|
616
|
+
self.release = asyncio.Event()
|
|
617
|
+
self.fail_with = None
|
|
618
|
+
self.cancelled = 0
|
|
619
|
+
|
|
620
|
+
async def lifespan(self):
|
|
621
|
+
self.state = {"text": "", "error": None}
|
|
622
|
+
yield
|
|
623
|
+
|
|
624
|
+
@command("run/steer")
|
|
625
|
+
async def run_steer(self, params):
|
|
626
|
+
self.create_task(self._run())
|
|
627
|
+
|
|
628
|
+
async def _run(self):
|
|
629
|
+
self.state["text"] += "a"
|
|
630
|
+
await self.release.wait()
|
|
631
|
+
if self.fail_with is not None:
|
|
632
|
+
self.state["error"] = self.fail_with
|
|
633
|
+
else:
|
|
634
|
+
self.state["text"] += "b"
|
|
635
|
+
|
|
636
|
+
async def on_assistant_transport_disconnect(self):
|
|
637
|
+
self.cancelled += 1
|
|
638
|
+
|
|
639
|
+
|
|
640
|
+
def post_background_chat(app):
|
|
641
|
+
return _PostStream(
|
|
642
|
+
app,
|
|
643
|
+
"/threads/t1/assistant-transport/api/chat",
|
|
644
|
+
{"commands": [{"type": "run/steer"}], "state": {}, "threadId": "t1"},
|
|
645
|
+
)
|
|
646
|
+
|
|
647
|
+
|
|
648
|
+
async def test_chat_stays_open_past_command_settle_and_eofs_at_idle():
|
|
649
|
+
async with statewire_client(BackgroundHarness) as (app, client):
|
|
650
|
+
async with post_background_chat(app) as stream:
|
|
651
|
+
assert stream.status == 200
|
|
652
|
+
buffer = b""
|
|
653
|
+
while b'"a"' not in buffer:
|
|
654
|
+
buffer += await stream.next_chunk()
|
|
655
|
+
with pytest.raises(TimeoutError):
|
|
656
|
+
await stream.next_chunk(0.1)
|
|
657
|
+
instance = (await app.state.pinned_host.directory.get("t1")).instance
|
|
658
|
+
instance.release.set()
|
|
659
|
+
buffer += await stream.until_closed()
|
|
660
|
+
frames = data_stream_frames(buffer.decode())
|
|
661
|
+
assert all(prefix == "aui-state" for prefix, _ in frames)
|
|
662
|
+
replica = apply_legacy_ops({}, state_ops(frames))
|
|
663
|
+
assert replica["text"] == "ab"
|
|
664
|
+
body = (
|
|
665
|
+
await client.post(
|
|
666
|
+
"/threads/t1/assistant-transport/api/status", json={"threadId": "t1"}
|
|
667
|
+
)
|
|
668
|
+
).json()
|
|
669
|
+
assert body["status"] == "completed"
|
|
670
|
+
|
|
671
|
+
|
|
672
|
+
async def test_run_error_recorded_into_state_is_an_error_frame_before_eof():
|
|
673
|
+
async with statewire_client(BackgroundHarness) as (app, client):
|
|
674
|
+
async with post_background_chat(app) as stream:
|
|
675
|
+
buffer = b""
|
|
676
|
+
while b'"a"' not in buffer:
|
|
677
|
+
buffer += await stream.next_chunk()
|
|
678
|
+
instance = (await app.state.pinned_host.directory.get("t1")).instance
|
|
679
|
+
instance.fail_with = {
|
|
680
|
+
"message": "payment required",
|
|
681
|
+
"reason": "payment-required",
|
|
682
|
+
}
|
|
683
|
+
instance.release.set()
|
|
684
|
+
buffer += await stream.until_closed()
|
|
685
|
+
frames = data_stream_frames(buffer.decode())
|
|
686
|
+
assert frames[-1] == ("3", "payment required")
|
|
687
|
+
replica = apply_legacy_ops({}, state_ops(frames))
|
|
688
|
+
assert replica["error"] == {
|
|
689
|
+
"message": "payment required",
|
|
690
|
+
"reason": "payment-required",
|
|
691
|
+
}
|
|
692
|
+
|
|
693
|
+
|
|
694
|
+
async def test_sse_mode_run_error_carries_the_structured_payload():
|
|
695
|
+
class SSEBackgroundHarness(BackgroundHarness):
|
|
696
|
+
assistant_transport_protocol = "assistant-transport"
|
|
697
|
+
|
|
698
|
+
async with statewire_client(SSEBackgroundHarness) as (app, client):
|
|
699
|
+
async with _PostStream(
|
|
700
|
+
app,
|
|
701
|
+
"/threads/t1/assistant-transport/api/chat",
|
|
702
|
+
{"commands": [{"type": "run/steer"}], "state": {}, "threadId": "t1"},
|
|
703
|
+
) as stream:
|
|
704
|
+
buffer = b""
|
|
705
|
+
while b'"a"' not in buffer:
|
|
706
|
+
buffer += await stream.next_chunk()
|
|
707
|
+
instance = (await app.state.pinned_host.directory.get("t1")).instance
|
|
708
|
+
instance.fail_with = {
|
|
709
|
+
"message": "payment required",
|
|
710
|
+
"reason": "payment-required",
|
|
711
|
+
}
|
|
712
|
+
instance.release.set()
|
|
713
|
+
buffer += await stream.until_closed()
|
|
714
|
+
frames = sse_frames(buffer.decode())
|
|
715
|
+
assert frames[-1] == "[DONE]"
|
|
716
|
+
assert frames[-2] == {
|
|
717
|
+
"type": "error",
|
|
718
|
+
"error": "payment required",
|
|
719
|
+
"payload": {"reason": "payment-required"},
|
|
720
|
+
}
|
|
721
|
+
|
|
722
|
+
|
|
723
|
+
async def test_disconnect_before_idle_fires_the_hook_and_the_run_continues():
|
|
724
|
+
async with statewire_client(BackgroundHarness) as (app, client):
|
|
725
|
+
async with post_background_chat(app) as stream:
|
|
726
|
+
buffer = b""
|
|
727
|
+
while b'"a"' not in buffer:
|
|
728
|
+
buffer += await stream.next_chunk()
|
|
729
|
+
await asyncio.sleep(0.05)
|
|
730
|
+
instance = (await app.state.pinned_host.directory.get("t1")).instance
|
|
731
|
+
assert instance.cancelled == 1
|
|
732
|
+
instance.release.set()
|
|
733
|
+
await asyncio.sleep(0.05)
|
|
734
|
+
response = await client.post(
|
|
735
|
+
"/threads/t1/assistant-transport/api/resume",
|
|
736
|
+
json={"commands": [], "state": {}, "threadId": "t1"},
|
|
737
|
+
)
|
|
738
|
+
assert response.headers["x-stream-status"] == "completed"
|
|
739
|
+
ops = state_ops(data_stream_frames(response.text))
|
|
740
|
+
assert ops[0]["value"]["text"] == "ab"
|
|
741
|
+
|
|
742
|
+
|
|
600
743
|
class RunHarness(AssistantTransport):
|
|
601
744
|
"""Host speaking the canonical run/* dialect; no legacy handlers."""
|
|
602
745
|
|
|
@@ -697,7 +840,8 @@ async def test_add_message_is_normalized_to_run_steer():
|
|
|
697
840
|
assert response.status_code == 200
|
|
698
841
|
instance = (await app.state.pinned_host.directory.get("t1")).instance
|
|
699
842
|
[(params, caller_id)] = instance.steered
|
|
700
|
-
assert set(params) == {"message"}
|
|
843
|
+
assert set(params) == {"message", "anchorMessageId"}
|
|
844
|
+
assert params["anchorMessageId"] is None
|
|
701
845
|
message = params["message"]
|
|
702
846
|
assert message["id"].startswith("legacy_")
|
|
703
847
|
assert message["role"] == "user"
|
|
@@ -951,7 +1095,9 @@ async def test_parent_id_without_commands_is_normalized_to_run_reload():
|
|
|
951
1095
|
frames = data_stream_frames(response.text)
|
|
952
1096
|
assert all(prefix == "aui-state" for prefix, _ in frames)
|
|
953
1097
|
instance = (await app.state.pinned_host.directory.get("t1")).instance
|
|
954
|
-
|
|
1098
|
+
[params] = instance.reloaded
|
|
1099
|
+
assert params["sourceId"] == "m2"
|
|
1100
|
+
assert params["messageId"].startswith("legacy_")
|
|
955
1101
|
|
|
956
1102
|
|
|
957
1103
|
async def test_reload_without_a_resolvable_child_is_an_error_frame():
|
|
@@ -992,6 +1138,59 @@ async def test_add_message_with_a_resolvable_parent_is_normalized_to_run_edit():
|
|
|
992
1138
|
assert params["message"]["parts"] == [{"type": "text", "text": "again"}]
|
|
993
1139
|
|
|
994
1140
|
|
|
1141
|
+
async def test_add_message_with_a_source_id_is_normalized_to_run_edit_with_it():
|
|
1142
|
+
async with statewire_client(AnchoredHarness) as (app, client):
|
|
1143
|
+
response = await post_run(
|
|
1144
|
+
client,
|
|
1145
|
+
{
|
|
1146
|
+
"commands": [
|
|
1147
|
+
{
|
|
1148
|
+
"type": "add-message",
|
|
1149
|
+
"message": {
|
|
1150
|
+
"role": "user",
|
|
1151
|
+
"parts": [{"type": "text", "text": "again"}],
|
|
1152
|
+
},
|
|
1153
|
+
"sourceId": "m9",
|
|
1154
|
+
}
|
|
1155
|
+
],
|
|
1156
|
+
"state": {},
|
|
1157
|
+
"parentId": "m1",
|
|
1158
|
+
},
|
|
1159
|
+
)
|
|
1160
|
+
assert response.status_code == 200
|
|
1161
|
+
instance = (await app.state.pinned_host.directory.get("t1")).instance
|
|
1162
|
+
assert instance.steered == []
|
|
1163
|
+
[params] = instance.edited
|
|
1164
|
+
assert params["sourceId"] == "m9"
|
|
1165
|
+
assert params["message"]["parts"] == [{"type": "text", "text": "again"}]
|
|
1166
|
+
|
|
1167
|
+
|
|
1168
|
+
async def test_add_message_with_a_null_source_id_and_no_parent_is_run_steer():
|
|
1169
|
+
async with statewire_client(AnchoredHarness) as (app, client):
|
|
1170
|
+
response = await post_run(
|
|
1171
|
+
client,
|
|
1172
|
+
{
|
|
1173
|
+
"commands": [
|
|
1174
|
+
{
|
|
1175
|
+
"type": "add-message",
|
|
1176
|
+
"message": {
|
|
1177
|
+
"role": "user",
|
|
1178
|
+
"parts": [{"type": "text", "text": "hi"}],
|
|
1179
|
+
},
|
|
1180
|
+
"sourceId": None,
|
|
1181
|
+
}
|
|
1182
|
+
],
|
|
1183
|
+
"state": {},
|
|
1184
|
+
},
|
|
1185
|
+
)
|
|
1186
|
+
assert response.status_code == 200
|
|
1187
|
+
instance = (await app.state.pinned_host.directory.get("t1")).instance
|
|
1188
|
+
assert instance.edited == []
|
|
1189
|
+
[(params, _)] = instance.steered
|
|
1190
|
+
assert set(params) == {"message", "anchorMessageId"}
|
|
1191
|
+
assert params["anchorMessageId"] is None
|
|
1192
|
+
|
|
1193
|
+
|
|
995
1194
|
async def test_add_message_with_an_unresolvable_parent_falls_back_to_run_steer():
|
|
996
1195
|
async with statewire_client(AnchoredHarness) as (app, client):
|
|
997
1196
|
response = await post_run(
|
|
@@ -1014,4 +1213,5 @@ async def test_add_message_with_an_unresolvable_parent_falls_back_to_run_steer()
|
|
|
1014
1213
|
instance = (await app.state.pinned_host.directory.get("t1")).instance
|
|
1015
1214
|
assert instance.edited == []
|
|
1016
1215
|
[(params, _)] = instance.steered
|
|
1017
|
-
assert set(params) == {"message"}
|
|
1216
|
+
assert set(params) == {"message", "anchorMessageId"}
|
|
1217
|
+
assert params["anchorMessageId"] == "unknown"
|
|
@@ -124,11 +124,11 @@ async def test_client_posts_the_legacy_body_and_applies_state_ops():
|
|
|
124
124
|
body={"system": "be nice"},
|
|
125
125
|
on_chunk=seen.append,
|
|
126
126
|
)
|
|
127
|
-
assert await client.send([{"type": "
|
|
127
|
+
assert await client.send([{"type": "demo/echo", "message": "hi"}]) is True
|
|
128
128
|
assert client.state == {"text": "hello"}
|
|
129
129
|
assert received == [
|
|
130
130
|
{
|
|
131
|
-
"commands": [{"type": "
|
|
131
|
+
"commands": [{"type": "demo/echo", "message": "hi"}],
|
|
132
132
|
"state": {"text": ""},
|
|
133
133
|
"threadId": "t1",
|
|
134
134
|
"system": "be nice",
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|