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.
Files changed (32) hide show
  1. {statewire-0.3.2 → statewire-0.4.0}/PKG-INFO +1 -1
  2. {statewire-0.3.2 → statewire-0.4.0}/pyproject.toml +1 -1
  3. {statewire-0.3.2 → statewire-0.4.0}/src/statewire/assistant_transport.py +148 -66
  4. {statewire-0.3.2 → statewire-0.4.0}/tests/test_assistant_transport.py +211 -11
  5. {statewire-0.3.2 → statewire-0.4.0}/tests/test_assistant_transport_client.py +2 -2
  6. {statewire-0.3.2 → statewire-0.4.0}/.gitignore +0 -0
  7. {statewire-0.3.2 → statewire-0.4.0}/README.md +0 -0
  8. {statewire-0.3.2 → statewire-0.4.0}/examples/__init__.py +0 -0
  9. {statewire-0.3.2 → statewire-0.4.0}/examples/demo_app.py +0 -0
  10. {statewire-0.3.2 → statewire-0.4.0}/src/statewire/__init__.py +0 -0
  11. {statewire-0.3.2 → statewire-0.4.0}/src/statewire/api.py +0 -0
  12. {statewire-0.3.2 → statewire-0.4.0}/src/statewire/assistant_transport_client.py +0 -0
  13. {statewire-0.3.2 → statewire-0.4.0}/src/statewire/client.py +0 -0
  14. {statewire-0.3.2 → statewire-0.4.0}/src/statewire/langgraph.py +0 -0
  15. {statewire-0.3.2 → statewire-0.4.0}/src/statewire/ops.py +0 -0
  16. {statewire-0.3.2 → statewire-0.4.0}/src/statewire/state.py +0 -0
  17. {statewire-0.3.2 → statewire-0.4.0}/tests/client_helpers.py +0 -0
  18. {statewire-0.3.2 → statewire-0.4.0}/tests/statewire_helpers.py +0 -0
  19. {statewire-0.3.2 → statewire-0.4.0}/tests/test_assistant_transport_facade.py +0 -0
  20. {statewire-0.3.2 → statewire-0.4.0}/tests/test_authorize.py +0 -0
  21. {statewire-0.3.2 → statewire-0.4.0}/tests/test_client.py +0 -0
  22. {statewire-0.3.2 → statewire-0.4.0}/tests/test_client_ws.py +0 -0
  23. {statewire-0.3.2 → statewire-0.4.0}/tests/test_commands.py +0 -0
  24. {statewire-0.3.2 → statewire-0.4.0}/tests/test_context.py +0 -0
  25. {statewire-0.3.2 → statewire-0.4.0}/tests/test_langgraph.py +0 -0
  26. {statewire-0.3.2 → statewire-0.4.0}/tests/test_lifespan.py +0 -0
  27. {statewire-0.3.2 → statewire-0.4.0}/tests/test_meta.py +0 -0
  28. {statewire-0.3.2 → statewire-0.4.0}/tests/test_state_proxy.py +0 -0
  29. {statewire-0.3.2 → statewire-0.4.0}/tests/test_statewire_hostable.py +0 -0
  30. {statewire-0.3.2 → statewire-0.4.0}/tests/test_stream.py +0 -0
  31. {statewire-0.3.2 → statewire-0.4.0}/tests/test_writer_lease.py +0 -0
  32. {statewire-0.3.2 → statewire-0.4.0}/tests/test_ws.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: statewire
3
- Version: 0.3.2
3
+ Version: 0.4.0
4
4
  Summary: Replicate one JSON object over an SSE op stream + a command endpoint
5
5
  Project-URL: Repository, https://github.com/assistant-ui/harness-sdk
6
6
  License-Expression: MIT
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "statewire"
3
- version = "0.3.2"
3
+ version = "0.4.0"
4
4
  description = "Replicate one JSON object over an SSE op stream + a command endpoint"
5
5
  readme = "README.md"
6
6
  license = "MIT"
@@ -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
- a fresh ``legacy_`` id is synthesized, image parts become
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. When the body carries a top-level ``parentId`` and
24
- ``get_message_child_id`` resolves it to a child, the command becomes
25
- ``run/edit`` with ``{"sourceId": <child>, "message": <wire message>}``.
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>}`` via ``get_message_child_id``
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
- state changes as legacy frames until every submitted command settles, then
38
- EOF; a rejection or crash becomes an error frame followed by EOF. Client
39
- disconnect before EOF invokes ``on_assistant_transport_disconnect`` (a plain
40
- overridable hook, not a command; default no-op); the run itself keeps
41
- executing server-side.
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 run settles. If no run ever started, 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
- after it settles, and ``{"isRunning": false, "status": "not_found",
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
- frame = json.dumps({"type": "error", "error": message})
124
- return f"data: {frame}\n\n"
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
- pass
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
- await self.get_message_child_id(parent_id)
420
- if parent_id is not None
421
- else None
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 {"method": "run/steer", "params": [{"message": message}]}
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 {"method": "run/reload", "params": [{"sourceId": source_id}]}
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
- while not (submit.done() and acked >= total and not open_pending):
491
- waiters = {q_get} if submit.done() else {q_get, submit}
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 not in done:
498
- continue
499
- item = q_get.result()
500
- if item is _LAGGARD:
501
- raise _RunError("state stream overflowed")
502
- data, finish = item
503
- envelope = json.loads(data)
504
- op_frames: list[str] = []
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.push(("error", run.error))
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 run settles.
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("message/add")
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 test_data_stream_response_streams_snapshot_then_ops_until_settle():
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": "message/add",
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": "message/add",
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": "message/add",
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": "message/add",
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": "message/add"}]},
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
- assert instance.reloaded == [{"sourceId": "m2"}]
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": "message/add", "message": "hi"}]) is True
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": "message/add", "message": "hi"}],
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