statewire 0.4.2__tar.gz → 0.4.4__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.4.2 → statewire-0.4.4}/PKG-INFO +1 -1
- {statewire-0.4.2 → statewire-0.4.4}/pyproject.toml +1 -1
- {statewire-0.4.2 → statewire-0.4.4}/src/statewire/api.py +143 -7
- {statewire-0.4.2 → statewire-0.4.4}/src/statewire/assistant_transport.py +50 -56
- statewire-0.4.4/src/statewire/testing.py +61 -0
- {statewire-0.4.2 → statewire-0.4.4}/tests/test_assistant_transport.py +41 -56
- {statewire-0.4.2 → statewire-0.4.4}/tests/test_authorize.py +73 -1
- statewire-0.4.4/tests/test_idle.py +91 -0
- statewire-0.4.4/tests/test_testing.py +64 -0
- {statewire-0.4.2 → statewire-0.4.4}/tests/test_ws.py +6 -3
- {statewire-0.4.2 → statewire-0.4.4}/.gitignore +0 -0
- {statewire-0.4.2 → statewire-0.4.4}/README.md +0 -0
- {statewire-0.4.2 → statewire-0.4.4}/examples/__init__.py +0 -0
- {statewire-0.4.2 → statewire-0.4.4}/examples/demo_app.py +0 -0
- {statewire-0.4.2 → statewire-0.4.4}/src/statewire/__init__.py +0 -0
- {statewire-0.4.2 → statewire-0.4.4}/src/statewire/assistant_transport_client.py +0 -0
- {statewire-0.4.2 → statewire-0.4.4}/src/statewire/client.py +0 -0
- {statewire-0.4.2 → statewire-0.4.4}/src/statewire/langgraph.py +0 -0
- {statewire-0.4.2 → statewire-0.4.4}/src/statewire/legacy.py +0 -0
- {statewire-0.4.2 → statewire-0.4.4}/src/statewire/ops.py +0 -0
- {statewire-0.4.2 → statewire-0.4.4}/src/statewire/state.py +0 -0
- {statewire-0.4.2 → statewire-0.4.4}/tests/client_helpers.py +0 -0
- {statewire-0.4.2 → statewire-0.4.4}/tests/statewire_helpers.py +0 -0
- {statewire-0.4.2 → statewire-0.4.4}/tests/test_assistant_transport_client.py +0 -0
- {statewire-0.4.2 → statewire-0.4.4}/tests/test_assistant_transport_facade.py +0 -0
- {statewire-0.4.2 → statewire-0.4.4}/tests/test_client.py +0 -0
- {statewire-0.4.2 → statewire-0.4.4}/tests/test_client_ws.py +0 -0
- {statewire-0.4.2 → statewire-0.4.4}/tests/test_commands.py +0 -0
- {statewire-0.4.2 → statewire-0.4.4}/tests/test_context.py +0 -0
- {statewire-0.4.2 → statewire-0.4.4}/tests/test_langgraph.py +0 -0
- {statewire-0.4.2 → statewire-0.4.4}/tests/test_legacy.py +0 -0
- {statewire-0.4.2 → statewire-0.4.4}/tests/test_lifespan.py +0 -0
- {statewire-0.4.2 → statewire-0.4.4}/tests/test_meta.py +0 -0
- {statewire-0.4.2 → statewire-0.4.4}/tests/test_state_proxy.py +0 -0
- {statewire-0.4.2 → statewire-0.4.4}/tests/test_statewire_hostable.py +0 -0
- {statewire-0.4.2 → statewire-0.4.4}/tests/test_stream.py +0 -0
- {statewire-0.4.2 → statewire-0.4.4}/tests/test_writer_lease.py +0 -0
|
@@ -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(
|
|
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:
|
|
@@ -29,7 +29,8 @@ the command whole):
|
|
|
29
29
|
- ``add-tool-result`` becomes ``run/input`` with ``{"requestId": <resolved>,
|
|
30
30
|
"response": {"output": <result>, "isError"?, "artifact"?}}``; the requestId
|
|
31
31
|
comes from ``get_input_request_id(toolCallId)``, whose default peeks the
|
|
32
|
-
unanswered tool-call request in ``state["inputRequests"]``
|
|
32
|
+
unanswered tool-call request in ``state["runs"][0]["inputRequests"]``
|
|
33
|
+
(None rejects),
|
|
33
34
|
and a ``modelContent`` field rejects (unsupported in ``run/input``).
|
|
34
35
|
- A body with a ``parentId`` and no commands is the legacy reload: it becomes
|
|
35
36
|
``run/reload`` with ``{"sourceId": <child>}`` via ``get_message_child_id``
|
|
@@ -39,25 +40,25 @@ A host without the target ``run/*`` handler rejects the batch via the
|
|
|
39
40
|
ordinary unknown-command path. The response streams state changes as legacy
|
|
40
41
|
frames until the instance is idle — every submitted command settled and no
|
|
41
42
|
live instance task (``create_task`` minus ``unref``) — then EOF. A run error
|
|
42
|
-
recorded into state (``get_run_error``; default
|
|
43
|
+
recorded into state (``get_run_error``; default
|
|
44
|
+
``state["runs"][0]["error"]``) becomes
|
|
43
45
|
an error frame before EOF, its fields beyond ``message`` carried as the
|
|
44
46
|
frame's structured payload where the format allows; a rejection or crash
|
|
45
|
-
becomes an error frame followed by EOF. Client disconnect before EOF
|
|
46
|
-
|
|
47
|
-
command; default no-op); the run itself keeps executing server-side.
|
|
47
|
+
becomes an error frame followed by EOF. Client disconnect before EOF
|
|
48
|
+
detaches the stream; the run keeps executing server-side.
|
|
48
49
|
|
|
49
50
|
``POST /assistant-transport/api/resume`` (same body shape, ``commands`` must
|
|
50
51
|
be absent or empty) reattaches to the run: a full-state ``set`` at the root,
|
|
51
52
|
then the live tail until the instance is idle. If no run ever started, the
|
|
52
53
|
response is 200 with an empty body and ``X-Stream-Status: not_found``. If the
|
|
53
54
|
run already completed, the replay carries ``X-Stream-Status: completed``. A
|
|
54
|
-
resume disconnect
|
|
55
|
+
resume disconnect detaches the same way.
|
|
55
56
|
|
|
56
|
-
``POST /assistant-transport/api/cancel``
|
|
57
|
-
|
|
58
|
-
|
|
59
|
-
|
|
60
|
-
|
|
57
|
+
``POST /assistant-transport/api/cancel`` dispatches the host's registered
|
|
58
|
+
``run/stop`` handler (if any; a rejection — e.g. wrong-state while idle — is
|
|
59
|
+
swallowed). A non-empty JSON-object body forwards as the ``run/stop``
|
|
60
|
+
params; an absent or empty body dispatches with none; any other body is a
|
|
61
|
+
400. It responds ``{"success": true}``.
|
|
61
62
|
|
|
62
63
|
``POST /assistant-transport/api/status`` responds 200 always:
|
|
63
64
|
``{"isRunning": true, "status": "running"}`` while a run is active,
|
|
@@ -80,7 +81,7 @@ import contextlib
|
|
|
80
81
|
import json
|
|
81
82
|
import secrets
|
|
82
83
|
import time
|
|
83
|
-
from typing import Any, AsyncIterator
|
|
84
|
+
from typing import Any, AsyncIterator
|
|
84
85
|
|
|
85
86
|
import httpx
|
|
86
87
|
from fastapi import FastAPI, HTTPException, Request
|
|
@@ -190,6 +191,11 @@ def _translate_op(replica: Any, op: dict[str, Any]) -> tuple[Any, str | None]:
|
|
|
190
191
|
_AT_CONSUMED_FIELDS = {"commands", "threadId", "state"}
|
|
191
192
|
|
|
192
193
|
|
|
194
|
+
def _run_entry(state: Any) -> Any:
|
|
195
|
+
runs = state.get("runs") if isinstance(state, dict) else None
|
|
196
|
+
return runs[0] if isinstance(runs, list) and runs else None
|
|
197
|
+
|
|
198
|
+
|
|
193
199
|
class _RunError(Exception):
|
|
194
200
|
def __init__(self, message: str, *, payload: Any = None) -> None:
|
|
195
201
|
super().__init__(message)
|
|
@@ -225,46 +231,30 @@ class AssistantTransport(Statewire):
|
|
|
225
231
|
|
|
226
232
|
_at_run: _Run | None = None
|
|
227
233
|
|
|
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
234
|
async def get_run_error(self) -> dict[str, Any] | None:
|
|
243
235
|
"""The run error to surface as the end-of-stream error frame.
|
|
244
236
|
|
|
245
237
|
The default peeks the RunManager root-state mount: a non-null
|
|
246
|
-
``state["error"]`` object with a string ``message`` is the
|
|
247
|
-
run error; its fields beyond ``message`` travel as the
|
|
248
|
-
structured payload where the format carries one. None ends
|
|
249
|
-
stream cleanly."""
|
|
250
|
-
|
|
251
|
-
error =
|
|
238
|
+
``state["runs"][0]["error"]`` object with a string ``message`` is the
|
|
239
|
+
recorded run error; its fields beyond ``message`` travel as the
|
|
240
|
+
frame's structured payload where the format carries one. None ends
|
|
241
|
+
the stream cleanly."""
|
|
242
|
+
entry = _run_entry(plain(self.state))
|
|
243
|
+
error = entry.get("error") if isinstance(entry, dict) else None
|
|
252
244
|
if isinstance(error, dict) and isinstance(error.get("message"), str):
|
|
253
245
|
return error
|
|
254
246
|
return None
|
|
255
247
|
|
|
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
248
|
async def get_input_request_id(self, tool_call_id: str) -> str | None:
|
|
260
249
|
"""Resolve a legacy toolCallId to the pending run/input requestId.
|
|
261
250
|
|
|
262
251
|
The default peeks the RunManager root-state mount:
|
|
263
|
-
``state["inputRequests"]`` holds ``{..., "response":
|
|
264
|
-
the unanswered tool-call request matching
|
|
265
|
-
rejects the legacy ``add-tool-result``
|
|
266
|
-
|
|
267
|
-
|
|
252
|
+
``state["runs"][0]["inputRequests"]`` holds ``{..., "response":
|
|
253
|
+
null}`` entries; the unanswered tool-call request matching
|
|
254
|
+
``tool_call_id`` wins. None rejects the legacy ``add-tool-result``
|
|
255
|
+
command."""
|
|
256
|
+
entry = _run_entry(plain(self.state))
|
|
257
|
+
requests = entry.get("inputRequests") if isinstance(entry, dict) else None
|
|
268
258
|
if not isinstance(requests, list):
|
|
269
259
|
return None
|
|
270
260
|
for request in requests:
|
|
@@ -330,7 +320,7 @@ class AssistantTransport(Statewire):
|
|
|
330
320
|
run, q, client_id, lease, commands, parent_id
|
|
331
321
|
)
|
|
332
322
|
)
|
|
333
|
-
self.
|
|
323
|
+
self._tasks.discard(watcher)
|
|
334
324
|
return StreamingResponse(
|
|
335
325
|
self._assistant_transport_attach(fmt, run, initial=True),
|
|
336
326
|
media_type="text/event-stream",
|
|
@@ -371,13 +361,20 @@ class AssistantTransport(Statewire):
|
|
|
371
361
|
@route.post("/assistant-transport/api/cancel")
|
|
372
362
|
async def assistant_transport_cancel(self, request: Request) -> dict[str, Any]:
|
|
373
363
|
identity = await self.authorize(request)
|
|
374
|
-
|
|
375
|
-
|
|
376
|
-
if
|
|
377
|
-
|
|
364
|
+
raw = await request.body()
|
|
365
|
+
params: list[Any] = []
|
|
366
|
+
if raw:
|
|
367
|
+
try:
|
|
368
|
+
body = json.loads(raw)
|
|
369
|
+
except (json.JSONDecodeError, UnicodeDecodeError):
|
|
370
|
+
raise HTTPException(status_code=400, detail="body must be JSON") from None
|
|
371
|
+
if not isinstance(body, dict):
|
|
372
|
+
raise HTTPException(status_code=400, detail="body must be a JSON object")
|
|
373
|
+
if body:
|
|
374
|
+
params = [body]
|
|
378
375
|
if "run/stop" in type(self)._statewire_commands:
|
|
379
|
-
await self._assistant_transport_dispatch("run/stop", identity)
|
|
380
|
-
return {"success": True
|
|
376
|
+
await self._assistant_transport_dispatch("run/stop", identity, params)
|
|
377
|
+
return {"success": True}
|
|
381
378
|
|
|
382
379
|
@route.post("/assistant-transport/api/status")
|
|
383
380
|
async def assistant_transport_status(self, request: Request) -> dict[str, Any]:
|
|
@@ -445,14 +442,16 @@ class AssistantTransport(Statewire):
|
|
|
445
442
|
return member
|
|
446
443
|
return {"method": kind, "params": [command]}
|
|
447
444
|
|
|
448
|
-
async def _assistant_transport_dispatch(
|
|
445
|
+
async def _assistant_transport_dispatch(
|
|
446
|
+
self, method: str, identity: Any, params: list[Any]
|
|
447
|
+
) -> None:
|
|
449
448
|
client_id = "at-" + secrets.token_urlsafe(9)
|
|
450
449
|
q, _ = self._register(client_id, identity)
|
|
451
450
|
try:
|
|
452
451
|
lease = self._clients[client_id].lease
|
|
453
452
|
assert lease is not None
|
|
454
453
|
await self._admit(
|
|
455
|
-
client_id, 1, {"commands": [{"method": method, "params":
|
|
454
|
+
client_id, 1, {"commands": [{"method": method, "params": params}]}, lease
|
|
456
455
|
)
|
|
457
456
|
finally:
|
|
458
457
|
self._subscribers.pop(q, None)
|
|
@@ -480,7 +479,7 @@ class AssistantTransport(Statewire):
|
|
|
480
479
|
submit.done()
|
|
481
480
|
and acked >= total
|
|
482
481
|
and not open_pending
|
|
483
|
-
and not self.
|
|
482
|
+
and not self._tasks
|
|
484
483
|
)
|
|
485
484
|
|
|
486
485
|
def consume(item: Any) -> bool:
|
|
@@ -530,7 +529,7 @@ class AssistantTransport(Statewire):
|
|
|
530
529
|
if idle():
|
|
531
530
|
break
|
|
532
531
|
continue
|
|
533
|
-
waiters = {q_get, *self.
|
|
532
|
+
waiters = {q_get, *self._tasks}
|
|
534
533
|
if not submit.done():
|
|
535
534
|
waiters.add(submit)
|
|
536
535
|
done, _ = await asyncio.wait(
|
|
@@ -566,7 +565,6 @@ class AssistantTransport(Statewire):
|
|
|
566
565
|
live = run if not run.done else None
|
|
567
566
|
if live is not None:
|
|
568
567
|
live.attachments.append(aq)
|
|
569
|
-
completed = False
|
|
570
568
|
try:
|
|
571
569
|
yield fmt.state([_legacy_op("set", [], run.replica)])
|
|
572
570
|
if live is not None:
|
|
@@ -588,14 +586,10 @@ class AssistantTransport(Statewire):
|
|
|
588
586
|
yield fmt.error(run.error, run.error_payload)
|
|
589
587
|
if fmt.end:
|
|
590
588
|
yield fmt.end
|
|
591
|
-
completed = True
|
|
592
589
|
finally:
|
|
593
590
|
if live is not None:
|
|
594
591
|
with contextlib.suppress(ValueError):
|
|
595
592
|
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
593
|
|
|
600
594
|
|
|
601
595
|
_HOP_HEADERS = {"host", "content-length", "transfer-encoding", "connection"}
|
|
@@ -0,0 +1,61 @@
|
|
|
1
|
+
"""Test helpers: real command contexts for driving @command handlers directly."""
|
|
2
|
+
|
|
3
|
+
import json
|
|
4
|
+
from typing import Any, Callable
|
|
5
|
+
|
|
6
|
+
from pinned.host import invalid_id_reason
|
|
7
|
+
|
|
8
|
+
from .api import (
|
|
9
|
+
_LAGGARD,
|
|
10
|
+
_UNSET,
|
|
11
|
+
Statewire,
|
|
12
|
+
StatewireClientHandle,
|
|
13
|
+
StatewireCommandContext,
|
|
14
|
+
_Invocation,
|
|
15
|
+
)
|
|
16
|
+
from .state import _frozen
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
def command_context(
|
|
20
|
+
sw: Statewire,
|
|
21
|
+
*,
|
|
22
|
+
client_id: str = "test",
|
|
23
|
+
identity: Any = None,
|
|
24
|
+
context: dict[str, Any] | None = None,
|
|
25
|
+
seq: int = 1,
|
|
26
|
+
) -> tuple[StatewireCommandContext, Callable[[], list[dict[str, Any]]]]:
|
|
27
|
+
"""A real ``StatewireCommandContext`` for calling a ``@command`` handler
|
|
28
|
+
directly, plus ``effects()``: each call flushes the host and returns every
|
|
29
|
+
envelope delivered to the client so far, decoded (attach snapshot excluded)."""
|
|
30
|
+
if not isinstance(sw, Statewire):
|
|
31
|
+
raise TypeError("command_context() requires a Statewire instance")
|
|
32
|
+
if sw._state is _UNSET:
|
|
33
|
+
raise RuntimeError(
|
|
34
|
+
"state is not set; assign sw.state (or enter lifespan) first"
|
|
35
|
+
)
|
|
36
|
+
reason = invalid_id_reason(client_id)
|
|
37
|
+
if reason is not None:
|
|
38
|
+
raise ValueError(f"invalid client_id: {reason}")
|
|
39
|
+
if seq < 1:
|
|
40
|
+
raise ValueError("seq must be >= 1")
|
|
41
|
+
if context is not None and not isinstance(context, dict):
|
|
42
|
+
raise TypeError("context must be a dict")
|
|
43
|
+
q, _snapshot = sw._register(client_id, identity)
|
|
44
|
+
client = sw._clients[client_id]
|
|
45
|
+
client.context = _frozen(context) if context is not None else None
|
|
46
|
+
inv = _Invocation(client, seq)
|
|
47
|
+
client.inflight[seq] = inv
|
|
48
|
+
ctx = StatewireCommandContext(StatewireClientHandle(sw, client), inv)
|
|
49
|
+
captured: list[dict[str, Any]] = []
|
|
50
|
+
|
|
51
|
+
def effects() -> list[dict[str, Any]]:
|
|
52
|
+
sw.flush()
|
|
53
|
+
while not q.empty():
|
|
54
|
+
item = q.get_nowait()
|
|
55
|
+
if item is _LAGGARD:
|
|
56
|
+
raise RuntimeError("test client fell behind its stream queue")
|
|
57
|
+
data, _finish = item
|
|
58
|
+
captured.append(json.loads(data))
|
|
59
|
+
return captured
|
|
60
|
+
|
|
61
|
+
return ctx, effects
|
|
@@ -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
|
|
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
|
|
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
|
|
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
|
|
558
|
-
assert instance.cancelled == 0
|
|
530
|
+
assert response.json() == {"success": True}
|
|
559
531
|
|
|
560
532
|
|
|
561
|
-
async def
|
|
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
|
|
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,10 +586,9 @@ 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
|
-
self.state = {"text": "", "
|
|
591
|
+
self.state = {"text": "", "runs": []}
|
|
622
592
|
yield
|
|
623
593
|
|
|
624
594
|
@command("run/steer")
|
|
@@ -629,13 +599,10 @@ class BackgroundHarness(AssistantTransport):
|
|
|
629
599
|
self.state["text"] += "a"
|
|
630
600
|
await self.release.wait()
|
|
631
601
|
if self.fail_with is not None:
|
|
632
|
-
self.state["
|
|
602
|
+
self.state["runs"] = [{"error": self.fail_with}]
|
|
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(
|
|
@@ -685,7 +652,7 @@ async def test_run_error_recorded_into_state_is_an_error_frame_before_eof():
|
|
|
685
652
|
frames = data_stream_frames(buffer.decode())
|
|
686
653
|
assert frames[-1] == ("3", "payment required")
|
|
687
654
|
replica = apply_legacy_ops({}, state_ops(frames))
|
|
688
|
-
assert replica["error"] == {
|
|
655
|
+
assert replica["runs"][0]["error"] == {
|
|
689
656
|
"message": "payment required",
|
|
690
657
|
"reason": "payment-required",
|
|
691
658
|
}
|
|
@@ -720,7 +687,7 @@ async def test_sse_mode_run_error_carries_the_structured_payload():
|
|
|
720
687
|
}
|
|
721
688
|
|
|
722
689
|
|
|
723
|
-
async def
|
|
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 =
|
|
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
|
|
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):
|
|
@@ -790,7 +752,7 @@ class InputHarness(RunHarness):
|
|
|
790
752
|
async def lifespan(self):
|
|
791
753
|
self.state = {
|
|
792
754
|
"messages": [],
|
|
793
|
-
"inputRequests": [
|
|
755
|
+
"runs": [{"inputRequests": [
|
|
794
756
|
{
|
|
795
757
|
"type": "tool-call",
|
|
796
758
|
"id": "req-old",
|
|
@@ -809,7 +771,7 @@ class InputHarness(RunHarness):
|
|
|
809
771
|
"toolCallId": "call-2",
|
|
810
772
|
"response": None,
|
|
811
773
|
},
|
|
812
|
-
],
|
|
774
|
+
]}],
|
|
813
775
|
}
|
|
814
776
|
yield
|
|
815
777
|
|
|
@@ -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
|
|
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 ==
|
|
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
|
|
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"]
|
|
@@ -0,0 +1,64 @@
|
|
|
1
|
+
import pytest
|
|
2
|
+
from statewire import Statewire, StatewireCommandContext, command, plain
|
|
3
|
+
from statewire.testing import command_context
|
|
4
|
+
from statewire_helpers import make_instance
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
class Board(Statewire):
|
|
8
|
+
async def lifespan(self):
|
|
9
|
+
self.state = {"posts": []}
|
|
10
|
+
yield
|
|
11
|
+
|
|
12
|
+
@command
|
|
13
|
+
async def post(self, text: str, *, ctx):
|
|
14
|
+
self.state["posts"].append({"text": text, "by": ctx.caller.identity})
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
def make_board():
|
|
18
|
+
sw = make_instance(Board)
|
|
19
|
+
sw.state = {"posts": []}
|
|
20
|
+
return sw
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
async def test_returns_real_context_backed_by_a_real_client():
|
|
24
|
+
sw = make_board()
|
|
25
|
+
ctx, _effects = command_context(
|
|
26
|
+
sw, client_id="c9", identity="alice", context={"tz": "utc"}
|
|
27
|
+
)
|
|
28
|
+
assert type(ctx) is StatewireCommandContext
|
|
29
|
+
assert ctx.caller.client_id == "c9"
|
|
30
|
+
assert ctx.caller.identity == "alice"
|
|
31
|
+
assert ctx.caller.context == {"tz": "utc"}
|
|
32
|
+
assert ctx.caller.is_connected
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
async def test_handler_sees_state_and_effects_capture_its_ops():
|
|
36
|
+
sw = make_board()
|
|
37
|
+
ctx, effects = command_context(sw, identity="alice")
|
|
38
|
+
await sw.post("hi", ctx=ctx)
|
|
39
|
+
assert plain(sw.state) == {"posts": [{"text": "hi", "by": "alice"}]}
|
|
40
|
+
(envelope,) = effects()
|
|
41
|
+
assert envelope["ops"] == [
|
|
42
|
+
{"op": "add", "path": ["posts", 0], "value": {"text": "hi", "by": "alice"}}
|
|
43
|
+
]
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
async def test_ack_is_captured_as_pending_res_and_ack_watermark():
|
|
47
|
+
sw = make_board()
|
|
48
|
+
ctx, effects = command_context(sw, seq=3)
|
|
49
|
+
ctx.ack()
|
|
50
|
+
assert effects() == [{"res": [{"seq": 3, "type": "pending"}], "ack": 3}]
|
|
51
|
+
|
|
52
|
+
|
|
53
|
+
async def test_rejects_invalid_input():
|
|
54
|
+
with pytest.raises(TypeError, match="Statewire instance"):
|
|
55
|
+
command_context(object()) # type: ignore[arg-type]
|
|
56
|
+
with pytest.raises(RuntimeError, match="state is not set"):
|
|
57
|
+
command_context(make_instance(Board))
|
|
58
|
+
sw = make_board()
|
|
59
|
+
with pytest.raises(ValueError, match="invalid client_id"):
|
|
60
|
+
command_context(sw, client_id="")
|
|
61
|
+
with pytest.raises(ValueError, match="seq must be >= 1"):
|
|
62
|
+
command_context(sw, seq=0)
|
|
63
|
+
with pytest.raises(TypeError, match="context must be a dict"):
|
|
64
|
+
command_context(sw, context="nope") # type: ignore[arg-type]
|
|
@@ -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
|
|
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
|
|
408
|
-
assert ws.
|
|
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
|
|
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
|