statewire 0.4.2__tar.gz → 0.4.3__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 (35) hide show
  1. {statewire-0.4.2 → statewire-0.4.3}/PKG-INFO +1 -1
  2. {statewire-0.4.2 → statewire-0.4.3}/pyproject.toml +1 -1
  3. {statewire-0.4.2 → statewire-0.4.3}/src/statewire/api.py +143 -7
  4. {statewire-0.4.2 → statewire-0.4.3}/src/statewire/assistant_transport.py +29 -43
  5. {statewire-0.4.2 → statewire-0.4.3}/tests/test_assistant_transport.py +36 -51
  6. {statewire-0.4.2 → statewire-0.4.3}/tests/test_authorize.py +73 -1
  7. statewire-0.4.3/tests/test_idle.py +91 -0
  8. {statewire-0.4.2 → statewire-0.4.3}/tests/test_ws.py +6 -3
  9. {statewire-0.4.2 → statewire-0.4.3}/.gitignore +0 -0
  10. {statewire-0.4.2 → statewire-0.4.3}/README.md +0 -0
  11. {statewire-0.4.2 → statewire-0.4.3}/examples/__init__.py +0 -0
  12. {statewire-0.4.2 → statewire-0.4.3}/examples/demo_app.py +0 -0
  13. {statewire-0.4.2 → statewire-0.4.3}/src/statewire/__init__.py +0 -0
  14. {statewire-0.4.2 → statewire-0.4.3}/src/statewire/assistant_transport_client.py +0 -0
  15. {statewire-0.4.2 → statewire-0.4.3}/src/statewire/client.py +0 -0
  16. {statewire-0.4.2 → statewire-0.4.3}/src/statewire/langgraph.py +0 -0
  17. {statewire-0.4.2 → statewire-0.4.3}/src/statewire/legacy.py +0 -0
  18. {statewire-0.4.2 → statewire-0.4.3}/src/statewire/ops.py +0 -0
  19. {statewire-0.4.2 → statewire-0.4.3}/src/statewire/state.py +0 -0
  20. {statewire-0.4.2 → statewire-0.4.3}/tests/client_helpers.py +0 -0
  21. {statewire-0.4.2 → statewire-0.4.3}/tests/statewire_helpers.py +0 -0
  22. {statewire-0.4.2 → statewire-0.4.3}/tests/test_assistant_transport_client.py +0 -0
  23. {statewire-0.4.2 → statewire-0.4.3}/tests/test_assistant_transport_facade.py +0 -0
  24. {statewire-0.4.2 → statewire-0.4.3}/tests/test_client.py +0 -0
  25. {statewire-0.4.2 → statewire-0.4.3}/tests/test_client_ws.py +0 -0
  26. {statewire-0.4.2 → statewire-0.4.3}/tests/test_commands.py +0 -0
  27. {statewire-0.4.2 → statewire-0.4.3}/tests/test_context.py +0 -0
  28. {statewire-0.4.2 → statewire-0.4.3}/tests/test_langgraph.py +0 -0
  29. {statewire-0.4.2 → statewire-0.4.3}/tests/test_legacy.py +0 -0
  30. {statewire-0.4.2 → statewire-0.4.3}/tests/test_lifespan.py +0 -0
  31. {statewire-0.4.2 → statewire-0.4.3}/tests/test_meta.py +0 -0
  32. {statewire-0.4.2 → statewire-0.4.3}/tests/test_state_proxy.py +0 -0
  33. {statewire-0.4.2 → statewire-0.4.3}/tests/test_statewire_hostable.py +0 -0
  34. {statewire-0.4.2 → statewire-0.4.3}/tests/test_stream.py +0 -0
  35. {statewire-0.4.2 → statewire-0.4.3}/tests/test_writer_lease.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: statewire
3
- Version: 0.4.2
3
+ Version: 0.4.3
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.4.2"
3
+ version = "0.4.3"
4
4
  description = "Replicate one JSON object over an SSE op stream + a command endpoint"
5
5
  readme = "README.md"
6
6
  license = "MIT"
@@ -88,7 +88,7 @@ import secrets
88
88
  import sys
89
89
  import time
90
90
  from dataclasses import dataclass, field
91
- from typing import Any, AsyncIterator, Awaitable, Callable
91
+ from typing import Any, AsyncIterator, Awaitable, Callable, Coroutine
92
92
 
93
93
  from fastapi import HTTPException, Request, WebSocket, WebSocketException
94
94
  from pinned import HostableContext, PinnedAPI
@@ -366,6 +366,33 @@ def _command_seq_of(request: Request) -> int:
366
366
  return seq
367
367
 
368
368
 
369
+ def _ws_frame_headers(raw: str) -> dict[str, str] | None:
370
+ """The optional first client frame: ``{"type": "headers", "headers": {...}}``.
371
+ Returns the map, ``None`` for any other frame; raises on a malformed one."""
372
+ try:
373
+ frame = json.loads(raw)
374
+ except (json.JSONDecodeError, UnicodeDecodeError):
375
+ return None
376
+ if not isinstance(frame, dict) or frame.get("type") != "headers":
377
+ return None
378
+ headers = frame.get("headers")
379
+ if not isinstance(headers, dict) or not all(
380
+ isinstance(k, str) and isinstance(v, str) for k, v in headers.items()
381
+ ):
382
+ raise ValueError("headers frame must carry a string map")
383
+ return headers
384
+
385
+
386
+ def _headers_connection(websocket: WebSocket, extra: dict[str, str]) -> HTTPConnection:
387
+ """The websocket's request view with the frame headers merged in (frame wins)."""
388
+ replaced = {name.lower().encode("latin-1") for name in extra}
389
+ merged = [(k, v) for k, v in websocket.scope["headers"] if k not in replaced]
390
+ merged += [
391
+ (k.lower().encode("latin-1"), v.encode("latin-1")) for k, v in extra.items()
392
+ ]
393
+ return HTTPConnection({**websocket.scope, "headers": merged})
394
+
395
+
369
396
  def _ws_frame_command(raw: str) -> tuple[dict[str, Any], int] | None:
370
397
  try:
371
398
  frame = json.loads(raw)
@@ -394,6 +421,7 @@ def _reserved_paths() -> set[str]:
394
421
  class Statewire(PinnedAPI):
395
422
  stream_queue_size: int = 256
396
423
  heartbeat_interval: float = 15.0
424
+ ws_init_timeout: float = 5.0
397
425
  client_record_limit: int = 1024
398
426
  envelope_ts: bool = False
399
427
  strip_invalid_types: bool = False
@@ -438,6 +466,7 @@ class Statewire(PinnedAPI):
438
466
  self._flush_mode = False
439
467
  self._draining = False
440
468
  self._drain_armed = False
469
+ self._tasks: set["asyncio.Task[Any]"] = set()
441
470
 
442
471
  async def lifespan(self) -> AsyncIterator[None]:
443
472
  if self._state is _UNSET:
@@ -585,6 +614,37 @@ class Statewire(PinnedAPI):
585
614
  finally:
586
615
  self._draining = False
587
616
 
617
+ # ─── Tasks and idleness ─────────────────────────────────
618
+
619
+ def create_task(self, coro: Coroutine[Any, Any, Any]) -> "asyncio.Task[Any]":
620
+ task = super().create_task(coro)
621
+ self._tasks.add(task)
622
+ task.add_done_callback(self._tasks.discard)
623
+ return task
624
+
625
+ def unref(self, task: "asyncio.Task[Any]") -> "asyncio.Task[Any]":
626
+ self._tasks.discard(task)
627
+ return super().unref(task)
628
+
629
+ async def idle(self) -> None:
630
+ """Wait until the instance is quiescent: every admitted command settled
631
+ and no tracked task (``create_task`` minus ``unref``) running."""
632
+ while True:
633
+ self.drain()
634
+ self.flush()
635
+ waiters = {
636
+ *self._tasks,
637
+ *(
638
+ inv.task
639
+ for client in self._clients.values()
640
+ for inv in client.inflight.values()
641
+ if inv.task is not None
642
+ ),
643
+ }
644
+ if not waiters:
645
+ return
646
+ await asyncio.wait(waiters, return_when=asyncio.FIRST_COMPLETED)
647
+
588
648
  # ─── Transient ops and finish ───────────────────────────
589
649
 
590
650
  def emit(self, value: dict[str, Any]) -> None:
@@ -1005,26 +1065,68 @@ class Statewire(PinnedAPI):
1005
1065
 
1006
1066
  @route.websocket("/ws")
1007
1067
  async def ws(self, websocket: WebSocket) -> None:
1008
- try:
1009
- identity = await self.authorize(websocket)
1010
- except HTTPException as exc:
1011
- reason = exc.detail if isinstance(exc.detail, str) else "unauthorized"
1012
- raise WebSocketException(code=1008, reason=reason) from None
1013
1068
  client_id = websocket.query_params.get("client", "")
1014
1069
  invalid = invalid_id_reason(client_id)
1015
1070
  if invalid is not None:
1016
1071
  raise WebSocketException(code=1008, reason=f"invalid client: {invalid}")
1072
+ try:
1073
+ identity = await self.authorize(websocket)
1074
+ authorized = True
1075
+ except HTTPException as exc:
1076
+ identity = _UNSET
1077
+ authorized = False
1078
+ denial = exc.detail if isinstance(exc.detail, str) else "unauthorized"
1017
1079
  offered = websocket.scope.get("subprotocols") or []
1018
1080
  await websocket.accept(
1019
1081
  subprotocol=ops.WS_SUBPROTOCOL if ops.WS_SUBPROTOCOL in offered else None
1020
1082
  )
1083
+ if not authorized:
1084
+ identity = await self._ws_init_auth(websocket, denial)
1085
+ if identity is _UNSET:
1086
+ return
1021
1087
  q, snapshot = self._register(client_id, identity)
1022
1088
  lease = self._clients[client_id].lease
1023
1089
  try:
1024
- await self._ws_pump(websocket, q, client_id, snapshot, lease)
1090
+ await self._ws_pump(
1091
+ websocket, q, client_id, snapshot, lease, allow_headers=authorized
1092
+ )
1025
1093
  finally:
1026
1094
  self._subscribers.pop(q, None)
1027
1095
 
1096
+ async def _ws_init_auth(self, websocket: WebSocket, denial: str) -> Any:
1097
+ """Unauthenticated attach: the headers frame must arrive within the init
1098
+ window; the snapshot is held until it is authorized. Returns the identity
1099
+ or ``_UNSET`` after closing the socket."""
1100
+ try:
1101
+ raw = await asyncio.wait_for(
1102
+ websocket.receive_text(), self.ws_init_timeout
1103
+ )
1104
+ except TimeoutError:
1105
+ await websocket.close(code=4401, reason=denial)
1106
+ return _UNSET
1107
+ except WebSocketDisconnect:
1108
+ return _UNSET
1109
+ except KeyError:
1110
+ await websocket.close(code=1008)
1111
+ return _UNSET
1112
+ try:
1113
+ headers = _ws_frame_headers(raw)
1114
+ connection = (
1115
+ None if headers is None else _headers_connection(websocket, headers)
1116
+ )
1117
+ except ValueError:
1118
+ await websocket.close(code=1008)
1119
+ return _UNSET
1120
+ if connection is None:
1121
+ await websocket.close(code=4401, reason=denial)
1122
+ return _UNSET
1123
+ try:
1124
+ return await self.authorize(connection)
1125
+ except HTTPException as exc:
1126
+ reason = exc.detail if isinstance(exc.detail, str) else "unauthorized"
1127
+ await websocket.close(code=4401, reason=reason)
1128
+ return _UNSET
1129
+
1028
1130
  async def _ws_pump(
1029
1131
  self,
1030
1132
  websocket: WebSocket,
@@ -1032,6 +1134,7 @@ class Statewire(PinnedAPI):
1032
1134
  client_id: str,
1033
1135
  snapshot: bytes,
1034
1136
  lease: str,
1137
+ allow_headers: bool = False,
1035
1138
  ) -> None:
1036
1139
  await websocket.send_text(snapshot.decode())
1037
1140
  q_task: asyncio.Task[Any] = asyncio.create_task(q.get())
@@ -1062,6 +1165,20 @@ class Statewire(PinnedAPI):
1062
1165
  except KeyError:
1063
1166
  await websocket.close(code=1008)
1064
1167
  return
1168
+ first = allow_headers
1169
+ allow_headers = False
1170
+ headers = None
1171
+ if first:
1172
+ try:
1173
+ headers = _ws_frame_headers(raw)
1174
+ except ValueError:
1175
+ await websocket.close(code=1008)
1176
+ return
1177
+ if headers is not None:
1178
+ if not await self._ws_restamp(websocket, client_id, headers):
1179
+ return
1180
+ io_task = asyncio.create_task(websocket.receive_text())
1181
+ continue
1065
1182
  parsed = _ws_frame_command(raw)
1066
1183
  if parsed is None:
1067
1184
  await websocket.close(code=1008)
@@ -1096,6 +1213,25 @@ class Statewire(PinnedAPI):
1096
1213
  q_task.cancel()
1097
1214
  io_task.cancel()
1098
1215
 
1216
+ async def _ws_restamp(
1217
+ self, websocket: WebSocket, client_id: str, headers: dict[str, str]
1218
+ ) -> bool:
1219
+ try:
1220
+ connection = _headers_connection(websocket, headers)
1221
+ except ValueError:
1222
+ await websocket.close(code=1008)
1223
+ return False
1224
+ try:
1225
+ identity = await self.authorize(connection)
1226
+ except HTTPException as exc:
1227
+ reason = exc.detail if isinstance(exc.detail, str) else "unauthorized"
1228
+ await websocket.close(code=4401, reason=reason)
1229
+ return False
1230
+ client = self._clients.get(client_id)
1231
+ if client is not None:
1232
+ client.identity = identity
1233
+ return True
1234
+
1099
1235
  async def _ws_execute(
1100
1236
  self, client_id: str, seq: int, command: dict[str, Any], lease: str
1101
1237
  ) -> dict[str, Any] | None:
@@ -42,22 +42,21 @@ live instance task (``create_task`` minus ``unref``) — then EOF. A run error
42
42
  recorded into state (``get_run_error``; default ``state["error"]``) becomes
43
43
  an error frame before EOF, its fields beyond ``message`` carried as the
44
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.
45
+ becomes an error frame followed by EOF. Client disconnect before EOF
46
+ detaches the stream; the run keeps executing server-side.
48
47
 
49
48
  ``POST /assistant-transport/api/resume`` (same body shape, ``commands`` must
50
49
  be absent or empty) reattaches to the run: a full-state ``set`` at the root,
51
50
  then the live tail until the instance is idle. If no run ever started, the
52
51
  response is 200 with an empty body and ``X-Stream-Status: not_found``. If the
53
52
  run already completed, the replay carries ``X-Stream-Status: completed``. A
54
- resume disconnect never invokes the cancel hook.
53
+ resume disconnect detaches the same way.
55
54
 
56
- ``POST /assistant-transport/api/cancel`` invokes the same disconnect hook,
57
- dispatches the host's registered ``run/stop`` handler (if any; a rejection —
58
- e.g. wrong-state while idle — is swallowed), and responds
59
- ``{"success": true, "found": <bool>}``; ``found`` is true when a live run
60
- existed at request time.
55
+ ``POST /assistant-transport/api/cancel`` dispatches the host's registered
56
+ ``run/stop`` handler (if any; a rejection — e.g. wrong-state while idle — is
57
+ swallowed). A non-empty JSON-object body forwards as the ``run/stop``
58
+ params; an absent or empty body dispatches with none; any other body is a
59
+ 400. It responds ``{"success": true}``.
61
60
 
62
61
  ``POST /assistant-transport/api/status`` responds 200 always:
63
62
  ``{"isRunning": true, "status": "running"}`` while a run is active,
@@ -80,7 +79,7 @@ import contextlib
80
79
  import json
81
80
  import secrets
82
81
  import time
83
- from typing import Any, AsyncIterator, Coroutine
82
+ from typing import Any, AsyncIterator
84
83
 
85
84
  import httpx
86
85
  from fastapi import FastAPI, HTTPException, Request
@@ -225,20 +224,6 @@ class AssistantTransport(Statewire):
225
224
 
226
225
  _at_run: _Run | None = None
227
226
 
228
- def __init__(self, ctx: Any) -> None:
229
- super().__init__(ctx)
230
- self._at_tasks: set["asyncio.Task[Any]"] = set()
231
-
232
- def create_task(self, coro: Coroutine[Any, Any, Any]) -> "asyncio.Task[Any]":
233
- task = super().create_task(coro)
234
- self._at_tasks.add(task)
235
- task.add_done_callback(self._at_tasks.discard)
236
- return task
237
-
238
- def unref(self, task: "asyncio.Task[Any]") -> "asyncio.Task[Any]":
239
- self._at_tasks.discard(task)
240
- return super().unref(task)
241
-
242
227
  async def get_run_error(self) -> dict[str, Any] | None:
243
228
  """The run error to surface as the end-of-stream error frame.
244
229
 
@@ -253,9 +238,6 @@ class AssistantTransport(Statewire):
253
238
  return error
254
239
  return None
255
240
 
256
- async def on_assistant_transport_disconnect(self) -> None:
257
- """The legacy client went away or asked to stop; treat as a cancel signal."""
258
-
259
241
  async def get_input_request_id(self, tool_call_id: str) -> str | None:
260
242
  """Resolve a legacy toolCallId to the pending run/input requestId.
261
243
 
@@ -330,7 +312,7 @@ class AssistantTransport(Statewire):
330
312
  run, q, client_id, lease, commands, parent_id
331
313
  )
332
314
  )
333
- self._at_tasks.discard(watcher)
315
+ self._tasks.discard(watcher)
334
316
  return StreamingResponse(
335
317
  self._assistant_transport_attach(fmt, run, initial=True),
336
318
  media_type="text/event-stream",
@@ -371,13 +353,20 @@ class AssistantTransport(Statewire):
371
353
  @route.post("/assistant-transport/api/cancel")
372
354
  async def assistant_transport_cancel(self, request: Request) -> dict[str, Any]:
373
355
  identity = await self.authorize(request)
374
- run = self._at_run
375
- found = run is not None and not run.done
376
- if found:
377
- await self.on_assistant_transport_disconnect()
356
+ raw = await request.body()
357
+ params: list[Any] = []
358
+ if raw:
359
+ try:
360
+ body = json.loads(raw)
361
+ except (json.JSONDecodeError, UnicodeDecodeError):
362
+ raise HTTPException(status_code=400, detail="body must be JSON") from None
363
+ if not isinstance(body, dict):
364
+ raise HTTPException(status_code=400, detail="body must be a JSON object")
365
+ if body:
366
+ params = [body]
378
367
  if "run/stop" in type(self)._statewire_commands:
379
- await self._assistant_transport_dispatch("run/stop", identity)
380
- return {"success": True, "found": found}
368
+ await self._assistant_transport_dispatch("run/stop", identity, params)
369
+ return {"success": True}
381
370
 
382
371
  @route.post("/assistant-transport/api/status")
383
372
  async def assistant_transport_status(self, request: Request) -> dict[str, Any]:
@@ -445,14 +434,16 @@ class AssistantTransport(Statewire):
445
434
  return member
446
435
  return {"method": kind, "params": [command]}
447
436
 
448
- async def _assistant_transport_dispatch(self, method: str, identity: Any) -> None:
437
+ async def _assistant_transport_dispatch(
438
+ self, method: str, identity: Any, params: list[Any]
439
+ ) -> None:
449
440
  client_id = "at-" + secrets.token_urlsafe(9)
450
441
  q, _ = self._register(client_id, identity)
451
442
  try:
452
443
  lease = self._clients[client_id].lease
453
444
  assert lease is not None
454
445
  await self._admit(
455
- client_id, 1, {"commands": [{"method": method, "params": []}]}, lease
446
+ client_id, 1, {"commands": [{"method": method, "params": params}]}, lease
456
447
  )
457
448
  finally:
458
449
  self._subscribers.pop(q, None)
@@ -480,7 +471,7 @@ class AssistantTransport(Statewire):
480
471
  submit.done()
481
472
  and acked >= total
482
473
  and not open_pending
483
- and not self._at_tasks
474
+ and not self._tasks
484
475
  )
485
476
 
486
477
  def consume(item: Any) -> bool:
@@ -530,7 +521,7 @@ class AssistantTransport(Statewire):
530
521
  if idle():
531
522
  break
532
523
  continue
533
- waiters = {q_get, *self._at_tasks}
524
+ waiters = {q_get, *self._tasks}
534
525
  if not submit.done():
535
526
  waiters.add(submit)
536
527
  done, _ = await asyncio.wait(
@@ -566,7 +557,6 @@ class AssistantTransport(Statewire):
566
557
  live = run if not run.done else None
567
558
  if live is not None:
568
559
  live.attachments.append(aq)
569
- completed = False
570
560
  try:
571
561
  yield fmt.state([_legacy_op("set", [], run.replica)])
572
562
  if live is not None:
@@ -588,14 +578,10 @@ class AssistantTransport(Statewire):
588
578
  yield fmt.error(run.error, run.error_payload)
589
579
  if fmt.end:
590
580
  yield fmt.end
591
- completed = True
592
581
  finally:
593
582
  if live is not None:
594
583
  with contextlib.suppress(ValueError):
595
584
  live.attachments.remove(aq)
596
- if not completed and initial:
597
- with contextlib.suppress(Exception):
598
- self.create_task(self.on_assistant_transport_disconnect())
599
585
 
600
586
 
601
587
  _HOP_HEADERS = {"host", "content-length", "transfer-encoding", "connection"}
@@ -21,7 +21,6 @@ class Harness(AssistantTransport):
21
21
  def __init__(self, ctx):
22
22
  super().__init__(ctx)
23
23
  self.seen = []
24
- self.cancelled = 0
25
24
 
26
25
  async def lifespan(self):
27
26
  self.state = {"messages": [], "text": "", "meta": {"title": None}}
@@ -55,9 +54,6 @@ class Harness(AssistantTransport):
55
54
  await asyncio.sleep(0.01)
56
55
  self.state["text"] += "b"
57
56
 
58
- async def on_assistant_transport_disconnect(self):
59
- self.cancelled += 1
60
-
61
57
 
62
58
  def data_stream_frames(text: str) -> list[tuple[str, object]]:
63
59
  frames = []
@@ -369,25 +365,6 @@ class _PostStream:
369
365
  await self._task
370
366
 
371
367
 
372
- async def test_client_disconnect_invokes_the_disconnect_hook():
373
- class SlowHarness(Harness):
374
- @command("slow")
375
- async def slow(self, cmd, *, ctx):
376
- ctx.ack()
377
- await asyncio.get_running_loop().create_future()
378
-
379
- async with statewire_client(SlowHarness) as (app, client):
380
- host = app.state.pinned_host
381
- async with _PostStream(
382
- app, "/threads/t1/assistant-transport/api/chat", {"commands": [{"type": "slow"}]}
383
- ) as stream:
384
- assert stream.status == 200
385
- assert (await stream.next_chunk()).startswith(b"aui-state:")
386
- await asyncio.sleep(0.05)
387
- instance = (await host.directory.get("t1")).instance
388
- assert instance.cancelled == 1
389
-
390
-
391
368
  class GatedHarness(Harness):
392
369
  def __init__(self, ctx):
393
370
  super().__init__(ctx)
@@ -495,7 +472,7 @@ async def test_mid_run_disconnect_then_resume_streams_snapshot_and_live_tail():
495
472
  assert replica["text"] == "ab"
496
473
 
497
474
 
498
- async def test_resume_disconnect_does_not_cancel_the_run():
475
+ async def test_resume_disconnect_leaves_the_run_executing():
499
476
  async with statewire_client(GatedHarness) as (app, client):
500
477
  host = app.state.pinned_host
501
478
  async with _PostStream(
@@ -505,13 +482,11 @@ async def test_resume_disconnect_does_not_cancel_the_run():
505
482
  await stream.next_chunk()
506
483
  await asyncio.sleep(0.05)
507
484
  instance = (await host.directory.get("t1")).instance
508
- cancelled_after_post = instance.cancelled
509
485
 
510
486
  async with post_resume(app) as resume:
511
487
  assert resume.status == 200
512
488
  await resume.next_chunk()
513
489
  await asyncio.sleep(0.05)
514
- assert instance.cancelled == cancelled_after_post
515
490
 
516
491
  instance.release.set()
517
492
  await asyncio.sleep(0.05)
@@ -540,25 +515,22 @@ async def test_sse_mode_resume_is_symmetric():
540
515
  assert frames[0]["operations"][0]["path"] == []
541
516
 
542
517
 
543
- async def test_cancel_without_a_live_run_reports_not_found():
518
+ async def test_cancel_without_a_live_run_succeeds():
544
519
  async with statewire_client(Harness) as (app, client):
545
520
  response = await client.post(
546
521
  "/threads/t1/assistant-transport/api/cancel", json={"threadId": "t1"}
547
522
  )
548
523
  assert response.status_code == 200
549
- assert response.json() == {"success": True, "found": False}
550
- instance = (await app.state.pinned_host.directory.get("t1")).instance
551
- assert instance.cancelled == 0
524
+ assert response.json() == {"success": True}
552
525
 
553
526
  await post_run(client, {"commands": [], "state": {}})
554
527
  response = await client.post(
555
528
  "/threads/t1/assistant-transport/api/cancel", json={"threadId": "t1"}
556
529
  )
557
- assert response.json() == {"success": True, "found": False}
558
- assert instance.cancelled == 0
530
+ assert response.json() == {"success": True}
559
531
 
560
532
 
561
- async def test_cancel_with_a_live_run_cancels_and_reports_found():
533
+ async def test_cancel_during_a_live_run_succeeds():
562
534
  async with statewire_client(GatedHarness) as (app, client):
563
535
  host = app.state.pinned_host
564
536
  async with _PostStream(
@@ -571,9 +543,8 @@ async def test_cancel_with_a_live_run_cancels_and_reports_found():
571
543
  response = await client.post(
572
544
  "/threads/t1/assistant-transport/api/cancel", json={"threadId": "t1"}
573
545
  )
574
- assert response.json() == {"success": True, "found": True}
546
+ assert response.json() == {"success": True}
575
547
  instance = (await host.directory.get("t1")).instance
576
- assert instance.cancelled == 1
577
548
  instance.release.set()
578
549
 
579
550
 
@@ -615,7 +586,6 @@ class BackgroundHarness(AssistantTransport):
615
586
  super().__init__(ctx)
616
587
  self.release = asyncio.Event()
617
588
  self.fail_with = None
618
- self.cancelled = 0
619
589
 
620
590
  async def lifespan(self):
621
591
  self.state = {"text": "", "error": None}
@@ -633,9 +603,6 @@ class BackgroundHarness(AssistantTransport):
633
603
  else:
634
604
  self.state["text"] += "b"
635
605
 
636
- async def on_assistant_transport_disconnect(self):
637
- self.cancelled += 1
638
-
639
606
 
640
607
  def post_background_chat(app):
641
608
  return _PostStream(
@@ -720,7 +687,7 @@ async def test_sse_mode_run_error_carries_the_structured_payload():
720
687
  }
721
688
 
722
689
 
723
- async def test_disconnect_before_idle_fires_the_hook_and_the_run_continues():
690
+ async def test_disconnect_before_idle_leaves_the_run_executing():
724
691
  async with statewire_client(BackgroundHarness) as (app, client):
725
692
  async with post_background_chat(app) as stream:
726
693
  buffer = b""
@@ -728,7 +695,6 @@ async def test_disconnect_before_idle_fires_the_hook_and_the_run_continues():
728
695
  buffer += await stream.next_chunk()
729
696
  await asyncio.sleep(0.05)
730
697
  instance = (await app.state.pinned_host.directory.get("t1")).instance
731
- assert instance.cancelled == 1
732
698
  instance.release.set()
733
699
  await asyncio.sleep(0.05)
734
700
  response = await client.post(
@@ -749,8 +715,7 @@ class RunHarness(AssistantTransport):
749
715
  self.edited = []
750
716
  self.inputs = []
751
717
  self.reloaded = []
752
- self.stopped = 0
753
- self.cancelled = 0
718
+ self.stopped = []
754
719
 
755
720
  async def lifespan(self):
756
721
  self.state = {"messages": []}
@@ -775,10 +740,7 @@ class RunHarness(AssistantTransport):
775
740
 
776
741
  @command("run/stop")
777
742
  async def run_stop(self, params=None, *, ctx):
778
- self.stopped += 1
779
-
780
- async def on_assistant_transport_disconnect(self):
781
- self.cancelled += 1
743
+ self.stopped.append(params)
782
744
 
783
745
 
784
746
  class AnchoredHarness(RunHarness):
@@ -938,10 +900,33 @@ async def test_cancel_dispatches_run_stop_when_registered():
938
900
  "/threads/t1/assistant-transport/api/cancel", json={"threadId": "t1"}
939
901
  )
940
902
  assert response.status_code == 200
941
- assert response.json() == {"success": True, "found": False}
903
+ assert response.json() == {"success": True}
904
+ instance = (await app.state.pinned_host.directory.get("t1")).instance
905
+ assert instance.stopped == [{"threadId": "t1"}]
906
+
907
+
908
+ async def test_cancel_forwards_the_body_as_run_stop_params():
909
+ async with statewire_client(RunHarness) as (app, client):
910
+ response = await client.post(
911
+ "/threads/t1/assistant-transport/api/cancel", json={"reason": "send_now"}
912
+ )
913
+ assert response.status_code == 200
914
+ instance = (await app.state.pinned_host.directory.get("t1")).instance
915
+ assert instance.stopped == [{"reason": "send_now"}]
916
+
917
+ response = await client.post("/threads/t1/assistant-transport/api/cancel")
918
+ assert response.status_code == 200
919
+ assert instance.stopped == [{"reason": "send_now"}, None]
920
+
921
+
922
+ async def test_cancel_rejects_a_non_object_body():
923
+ async with statewire_client(RunHarness) as (app, client):
924
+ response = await client.post(
925
+ "/threads/t1/assistant-transport/api/cancel", json=[1]
926
+ )
927
+ assert response.status_code == 400
942
928
  instance = (await app.state.pinned_host.directory.get("t1")).instance
943
- assert instance.stopped == 1
944
- assert instance.cancelled == 0
929
+ assert instance.stopped == []
945
930
 
946
931
 
947
932
  async def test_cancel_swallows_a_run_stop_rejection():
@@ -955,7 +940,7 @@ async def test_cancel_swallows_a_run_stop_rejection():
955
940
  "/threads/t1/assistant-transport/api/cancel", json={"threadId": "t1"}
956
941
  )
957
942
  assert response.status_code == 200
958
- assert response.json() == {"success": True, "found": False}
943
+ assert response.json() == {"success": True}
959
944
 
960
945
 
961
946
  async def test_add_tool_result_is_normalized_to_run_input():
@@ -201,6 +201,78 @@ async def test_identity_is_not_settable_via_context_updates():
201
201
  assert res["payload"] == {"identity": "alice"}
202
202
 
203
203
 
204
+ async def test_ws_headers_frame_authorizes_and_releases_the_snapshot():
205
+ async with statewire_client(Guarded) as (app, client):
206
+ async with ws_of(app) as ws:
207
+ assert ws.accepted
208
+ await ws.send_frame({"type": "headers", "headers": {"x-token": "secret"}})
209
+ frame = await ws.next_frame()
210
+ assert "ops" in frame and "syn" in frame
211
+ assert (await ws.command({"method": "inc", "params": []}))["ack"] == 1
212
+
213
+
214
+ async def test_ws_unauthorized_headers_frame_closes_4401():
215
+ async with statewire_client(Guarded) as (app, client):
216
+ async with ws_of(app) as ws:
217
+ assert ws.accepted
218
+ await ws.send_frame({"type": "headers", "headers": {"x-token": "wrong"}})
219
+ assert await ws.next_frame() is None
220
+ assert ws.close_code == 4401
221
+ assert ws.close_reason == "unauthorized"
222
+ assert instances[-1]._subscribers == {}
223
+
224
+
225
+ async def test_ws_init_timeout_closes_unauthenticated_sockets_4401():
226
+ class Quick(Guarded):
227
+ ws_init_timeout = 0.05
228
+
229
+ async with statewire_client(Quick) as (app, client):
230
+ async with ws_of(app) as ws:
231
+ assert ws.accepted
232
+ assert await ws.next_frame() is None
233
+ assert ws.close_code == 4401
234
+
235
+
236
+ async def test_ws_non_headers_first_frame_on_a_guarded_host_closes_4401():
237
+ async with statewire_client(Guarded) as (app, client):
238
+ async with ws_of(app) as ws:
239
+ assert ws.accepted
240
+ await ws.send_frame({"method": "inc", "params": [], "seq": 1})
241
+ assert await ws.next_frame() is None
242
+ assert ws.close_code == 4401
243
+
244
+
245
+ async def test_ws_headers_frame_identity_reaches_ctx_caller():
246
+ async with statewire_client(Identified) as (app, client):
247
+ async with ws_of(app) as ws:
248
+ await ws.send_frame({"type": "headers", "headers": {"x-user": "wanda"}})
249
+ await ws.next_frame()
250
+ env = await ws.command({"method": "whoami", "params": []})
251
+ assert env["res"] == [
252
+ {"seq": 1, "type": "accepted", "payload": {"identity": "wanda"}}
253
+ ]
254
+
255
+
256
+ async def test_ws_malformed_headers_frame_closes_1008():
257
+ async with statewire_client(Identified) as (app, client):
258
+ async with ws_of(app) as ws:
259
+ await ws.next_frame()
260
+ await ws.send_frame({"type": "headers", "headers": {"a": 1}})
261
+ assert await ws.next_frame() is None
262
+ assert ws.close_code == 1008
263
+
264
+
265
+ async def test_ws_without_headers_frame_attaches_as_before_when_unguarded():
266
+ async with statewire_client(Identified) as (app, client):
267
+ async with ws_of(app) as ws:
268
+ assert ws.accepted
269
+ assert "syn" in (await ws.next_frame())
270
+ env = await ws.command({"method": "whoami", "params": []})
271
+ assert env["res"] == [
272
+ {"seq": 1, "type": "accepted", "payload": {"identity": None}}
273
+ ]
274
+
275
+
204
276
  class LegacyIdentified(AssistantTransport):
205
277
  seen = None
206
278
  stopped = None
@@ -217,7 +289,7 @@ class LegacyIdentified(AssistantTransport):
217
289
  self.seen = ctx.caller.identity
218
290
 
219
291
  @command("run/stop")
220
- async def stop(self, *, ctx):
292
+ async def stop(self, params=None, *, ctx):
221
293
  self.stopped = ctx.caller.identity
222
294
 
223
295
 
@@ -0,0 +1,91 @@
1
+ import asyncio
2
+
3
+ from statewire_helpers import attach, post_command, statewire_client
4
+
5
+ from statewire import Statewire, command
6
+
7
+
8
+ def make_host():
9
+ class Host(Statewire):
10
+ def __init__(self, ctx):
11
+ super().__init__(ctx)
12
+ self.gates = [asyncio.Event(), asyncio.Event(), asyncio.Event()]
13
+
14
+ async def lifespan(self):
15
+ self.state = {"log": []}
16
+ yield
17
+
18
+ @command
19
+ async def hold(self, *, ctx):
20
+ ctx.ack()
21
+ await self.gates[0].wait()
22
+ self.state["log"].append("held")
23
+
24
+ @command
25
+ async def chain(self, *, ctx):
26
+ ctx.ack()
27
+ await self.gates[0].wait()
28
+ self.create_task(self._stage_a())
29
+
30
+ async def _stage_a(self):
31
+ await self.gates[1].wait()
32
+ self.create_task(self._stage_b())
33
+
34
+ async def _stage_b(self):
35
+ await self.gates[2].wait()
36
+ self.state["log"].append("done")
37
+
38
+ return Host
39
+
40
+
41
+ async def _instance_of(app):
42
+ return (await app.state.pinned_host.directory.get("t1")).instance
43
+
44
+
45
+ async def _settles(idle: "asyncio.Task[None]") -> bool:
46
+ await asyncio.sleep(0.05)
47
+ return idle.done()
48
+
49
+
50
+ async def test_idle_returns_promptly_on_a_fresh_instance():
51
+ async with statewire_client(make_host()) as (app, client):
52
+ await attach(app)
53
+ instance = await _instance_of(app)
54
+ await asyncio.wait_for(instance.idle(), timeout=1)
55
+
56
+
57
+ async def test_idle_blocks_on_a_pending_command_until_settle():
58
+ async with statewire_client(make_host()) as (app, client):
59
+ await attach(app)
60
+ assert (await post_command(client, {"method": "hold", "params": []}, seq=1)).status_code == 200
61
+ instance = await _instance_of(app)
62
+ idle = asyncio.create_task(instance.idle())
63
+ assert not await _settles(idle)
64
+ instance.gates[0].set()
65
+ await asyncio.wait_for(idle, timeout=1)
66
+ assert list(instance.state["log"]) == ["held"]
67
+
68
+
69
+ async def test_idle_blocks_on_a_tracked_task_until_it_finishes():
70
+ async with statewire_client(make_host()) as (app, client):
71
+ await attach(app)
72
+ instance = await _instance_of(app)
73
+ instance.create_task(instance._stage_b())
74
+ idle = asyncio.create_task(instance.idle())
75
+ assert not await _settles(idle)
76
+ instance.gates[2].set()
77
+ await asyncio.wait_for(idle, timeout=1)
78
+ assert list(instance.state["log"]) == ["done"]
79
+
80
+
81
+ async def test_one_idle_call_spans_a_command_task_task_chain():
82
+ async with statewire_client(make_host()) as (app, client):
83
+ await attach(app)
84
+ assert (await post_command(client, {"method": "chain", "params": []}, seq=1)).status_code == 200
85
+ instance = await _instance_of(app)
86
+ idle = asyncio.create_task(instance.idle())
87
+ for gate in instance.gates:
88
+ assert not await _settles(idle)
89
+ gate.set()
90
+ await asyncio.wait_for(idle, timeout=1)
91
+ assert list(instance.state["log"]) == ["done"]
@@ -396,16 +396,19 @@ async def test_gone_finish_reaches_ws_then_close():
396
396
  assert await ws.next_frame() is None
397
397
 
398
398
 
399
- async def test_authorize_rejection_closes_1008_and_registers_nothing():
399
+ async def test_authorize_rejection_closes_4401_and_registers_nothing():
400
400
  class Guarded(make_counter()):
401
+ ws_init_timeout = 0.05
402
+
401
403
  async def authorize(self, request: Request) -> None:
402
404
  if request.headers.get("x-token") != "secret":
403
405
  raise HTTPException(status_code=401, detail="unauthorized")
404
406
 
405
407
  async with statewire_client(Guarded) as (app, client):
406
408
  async with ws_of(app) as ws:
407
- assert ws.accepted is False
408
- assert ws.close_code == 1008
409
+ assert ws.accepted
410
+ assert await ws.next_frame() is None
411
+ assert ws.close_code == 4401
409
412
  assert ws.close_reason == "unauthorized"
410
413
  host = app.state.pinned_host
411
414
  record = await host.directory.get("t1")
File without changes
File without changes
File without changes