statewire 0.3.1__tar.gz → 0.3.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.
- {statewire-0.3.1 → statewire-0.3.3}/PKG-INFO +1 -1
- {statewire-0.3.1 → statewire-0.3.3}/pyproject.toml +1 -1
- {statewire-0.3.1 → statewire-0.3.3}/src/statewire/api.py +50 -14
- {statewire-0.3.1 → statewire-0.3.3}/src/statewire/assistant_transport.py +139 -66
- {statewire-0.3.1 → statewire-0.3.3}/tests/test_assistant_transport.py +197 -2
- statewire-0.3.3/tests/test_authorize.py +256 -0
- {statewire-0.3.1 → statewire-0.3.3}/tests/test_state_proxy.py +1 -1
- {statewire-0.3.1 → statewire-0.3.3}/tests/test_writer_lease.py +1 -1
- statewire-0.3.1/tests/test_authorize.py +0 -56
- {statewire-0.3.1 → statewire-0.3.3}/.gitignore +0 -0
- {statewire-0.3.1 → statewire-0.3.3}/README.md +0 -0
- {statewire-0.3.1 → statewire-0.3.3}/examples/__init__.py +0 -0
- {statewire-0.3.1 → statewire-0.3.3}/examples/demo_app.py +0 -0
- {statewire-0.3.1 → statewire-0.3.3}/src/statewire/__init__.py +0 -0
- {statewire-0.3.1 → statewire-0.3.3}/src/statewire/assistant_transport_client.py +0 -0
- {statewire-0.3.1 → statewire-0.3.3}/src/statewire/client.py +0 -0
- {statewire-0.3.1 → statewire-0.3.3}/src/statewire/langgraph.py +0 -0
- {statewire-0.3.1 → statewire-0.3.3}/src/statewire/ops.py +0 -0
- {statewire-0.3.1 → statewire-0.3.3}/src/statewire/state.py +0 -0
- {statewire-0.3.1 → statewire-0.3.3}/tests/client_helpers.py +0 -0
- {statewire-0.3.1 → statewire-0.3.3}/tests/statewire_helpers.py +0 -0
- {statewire-0.3.1 → statewire-0.3.3}/tests/test_assistant_transport_client.py +0 -0
- {statewire-0.3.1 → statewire-0.3.3}/tests/test_assistant_transport_facade.py +0 -0
- {statewire-0.3.1 → statewire-0.3.3}/tests/test_client.py +0 -0
- {statewire-0.3.1 → statewire-0.3.3}/tests/test_client_ws.py +0 -0
- {statewire-0.3.1 → statewire-0.3.3}/tests/test_commands.py +0 -0
- {statewire-0.3.1 → statewire-0.3.3}/tests/test_context.py +0 -0
- {statewire-0.3.1 → statewire-0.3.3}/tests/test_langgraph.py +0 -0
- {statewire-0.3.1 → statewire-0.3.3}/tests/test_lifespan.py +0 -0
- {statewire-0.3.1 → statewire-0.3.3}/tests/test_meta.py +0 -0
- {statewire-0.3.1 → statewire-0.3.3}/tests/test_statewire_hostable.py +0 -0
- {statewire-0.3.1 → statewire-0.3.3}/tests/test_stream.py +0 -0
- {statewire-0.3.1 → statewire-0.3.3}/tests/test_ws.py +0 -0
|
@@ -36,6 +36,16 @@ body over HTTP, as a ``res`` on the stream over WS.
|
|
|
36
36
|
Handlers that declare a keyword-only ``ctx`` parameter receive a
|
|
37
37
|
``StatewireCommandContext``: ``ctx.caller`` is the submitting client's
|
|
38
38
|
handle (holdable across a long run), ``ctx.ack()`` the invocation ack.
|
|
39
|
+
``authorize`` gates every ingress and returns the caller's identity
|
|
40
|
+
(``None`` = anonymous); the server stamps it on the client record per
|
|
41
|
+
attach and per admitted command submission — never client-supplied,
|
|
42
|
+
never on a rejected ingress — and drops it with the record. Handlers
|
|
43
|
+
read it as ``ctx.caller.identity``: the current session owner, live —
|
|
44
|
+
a newer authorized attach re-stamps it, so per-command trust decisions
|
|
45
|
+
read it before the first await. A client id is a within-thread session
|
|
46
|
+
id, not an authorization boundary: ``authorize`` must gate the thread
|
|
47
|
+
itself, since any authorized principal knowing a client id can attach,
|
|
48
|
+
receive the snapshot, and supersede that session.
|
|
39
49
|
|
|
40
50
|
The host also owns a scheduler: ``self.schedule(fn)`` runs ``fn`` at the
|
|
41
51
|
next drain (duplicate schedules coalesce), ``self.drain()`` drains now.
|
|
@@ -161,6 +171,7 @@ class _Client:
|
|
|
161
171
|
self.results: dict[int, tuple[float, dict[str, Any]]] = {}
|
|
162
172
|
self.inflight: dict[int, _Invocation] = {}
|
|
163
173
|
self.context: dict[str, Any] | None = None
|
|
174
|
+
self.identity: Any = None
|
|
164
175
|
|
|
165
176
|
def settle(self, seq: int, response: dict[str, Any] | None = None) -> None:
|
|
166
177
|
full = {"seq": seq, **(response if response is not None else {"type": "accepted"})}
|
|
@@ -207,8 +218,9 @@ class _Client:
|
|
|
207
218
|
class StatewireClientHandle:
|
|
208
219
|
"""Read-only view of the submitting client, reachable as ``ctx.caller``
|
|
209
220
|
in a ``@command`` handler. Holdable across a long run: ``context`` reads
|
|
210
|
-
the client's current stored context live, ``
|
|
211
|
-
|
|
221
|
+
the client's current stored context live, ``identity`` the current
|
|
222
|
+
session owner's identity, and ``is_connected`` reports whether the
|
|
223
|
+
client has an active stream attachment."""
|
|
212
224
|
|
|
213
225
|
def __init__(self, host: "Statewire", client: _Client) -> None:
|
|
214
226
|
self._host = host
|
|
@@ -222,6 +234,13 @@ class StatewireClientHandle:
|
|
|
222
234
|
def context(self) -> dict[str, Any] | None:
|
|
223
235
|
return self._client.context
|
|
224
236
|
|
|
237
|
+
@property
|
|
238
|
+
def identity(self) -> Any:
|
|
239
|
+
"""Current session owner's identity — live, re-stamped by a newer
|
|
240
|
+
authorized attach; per-command trust decisions read it before the
|
|
241
|
+
handler's first await."""
|
|
242
|
+
return self._client.identity
|
|
243
|
+
|
|
225
244
|
@property
|
|
226
245
|
def is_connected(self) -> bool:
|
|
227
246
|
client = self._client
|
|
@@ -592,14 +611,22 @@ class Statewire(PinnedAPI):
|
|
|
592
611
|
|
|
593
612
|
# ─── GET /stream ────────────────────────────────────────
|
|
594
613
|
|
|
595
|
-
async def authorize(self, request: HTTPConnection) ->
|
|
596
|
-
"""Auth seam: override to gate /stream, /commands and /ws; raise
|
|
614
|
+
async def authorize(self, request: HTTPConnection) -> Any:
|
|
615
|
+
"""Auth seam: override to gate /stream, /commands and /ws; raise
|
|
616
|
+
HTTPException to reject. The return value is the caller's identity —
|
|
617
|
+
any host-chosen object, ``None`` for anonymous — stamped on the client
|
|
618
|
+
record at each attach and each admitted command submission (fresh
|
|
619
|
+
identity wins) and read by handlers via ``ctx.caller.identity``."""
|
|
620
|
+
return None
|
|
597
621
|
|
|
598
|
-
def _register(
|
|
622
|
+
def _register(
|
|
623
|
+
self, client_id: str, identity: Any
|
|
624
|
+
) -> tuple[asyncio.Queue[Any], bytes]:
|
|
599
625
|
client = self._clients.get(client_id)
|
|
600
626
|
last = client.last_seq if client is not None else -1
|
|
601
627
|
if client is None:
|
|
602
628
|
client = self._clients[client_id] = _Client(self, client_id)
|
|
629
|
+
client.identity = identity
|
|
603
630
|
client.touched = time.monotonic()
|
|
604
631
|
self._flush()
|
|
605
632
|
self._evict_streams(
|
|
@@ -653,9 +680,9 @@ class Statewire(PinnedAPI):
|
|
|
653
680
|
|
|
654
681
|
@route.get("/stream")
|
|
655
682
|
async def stream(self, request: Request) -> Response:
|
|
656
|
-
await self.authorize(request)
|
|
683
|
+
identity = await self.authorize(request)
|
|
657
684
|
client_id = _client_id_of(request)
|
|
658
|
-
q, snapshot = self._register(client_id)
|
|
685
|
+
q, snapshot = self._register(client_id, identity)
|
|
659
686
|
return StreamingResponse(
|
|
660
687
|
self._stream_events(q, snapshot),
|
|
661
688
|
media_type="text/event-stream",
|
|
@@ -710,7 +737,7 @@ class Statewire(PinnedAPI):
|
|
|
710
737
|
|
|
711
738
|
@route.post("/commands")
|
|
712
739
|
async def commands(self, request: Request) -> Any:
|
|
713
|
-
await self.authorize(request)
|
|
740
|
+
identity = await self.authorize(request)
|
|
714
741
|
client_id = _client_id_of(request)
|
|
715
742
|
seq = _command_seq_of(request)
|
|
716
743
|
lease = request.headers.get(LEASE_HEADER)
|
|
@@ -719,7 +746,7 @@ class Statewire(PinnedAPI):
|
|
|
719
746
|
body = json.loads(raw)
|
|
720
747
|
except (json.JSONDecodeError, UnicodeDecodeError):
|
|
721
748
|
raise HTTPException(status_code=400, detail="body must be JSON") from None
|
|
722
|
-
return await self._submit(client_id, seq, body, lease)
|
|
749
|
+
return await self._submit(client_id, seq, body, lease, identity)
|
|
723
750
|
|
|
724
751
|
def _validate_submission(
|
|
725
752
|
self, body: Any
|
|
@@ -788,7 +815,12 @@ class Statewire(PinnedAPI):
|
|
|
788
815
|
return _Prepared(method, handler, params, wants_ctx), None
|
|
789
816
|
|
|
790
817
|
async def _admit(
|
|
791
|
-
self,
|
|
818
|
+
self,
|
|
819
|
+
client_id: str,
|
|
820
|
+
seq: int,
|
|
821
|
+
body: Any,
|
|
822
|
+
lease: str | None,
|
|
823
|
+
identity: Any = _UNSET,
|
|
792
824
|
) -> tuple[str, Any]:
|
|
793
825
|
client = self._clients.get(client_id)
|
|
794
826
|
if client is None:
|
|
@@ -803,6 +835,10 @@ class Statewire(PinnedAPI):
|
|
|
803
835
|
context, members, rejections = self._validate_submission(body)
|
|
804
836
|
if rejections:
|
|
805
837
|
return "rejected", rejections
|
|
838
|
+
# Only an admitted submission may re-stamp; _UNSET keeps the
|
|
839
|
+
# attach-time identity (WS frames, facade dispatch).
|
|
840
|
+
if identity is not _UNSET:
|
|
841
|
+
client.identity = identity
|
|
806
842
|
if context is not _UNSET:
|
|
807
843
|
client.context = context
|
|
808
844
|
invocations = self._start_batch(client, seq, members)
|
|
@@ -814,9 +850,9 @@ class Statewire(PinnedAPI):
|
|
|
814
850
|
return "executed", None
|
|
815
851
|
|
|
816
852
|
async def _submit(
|
|
817
|
-
self, client_id: str, seq: int, body: Any, lease: str | None
|
|
853
|
+
self, client_id: str, seq: int, body: Any, lease: str | None, identity: Any
|
|
818
854
|
) -> Any:
|
|
819
|
-
outcome, detail = await self._admit(client_id, seq, body, lease)
|
|
855
|
+
outcome, detail = await self._admit(client_id, seq, body, lease, identity)
|
|
820
856
|
match outcome:
|
|
821
857
|
case "unknown-client":
|
|
822
858
|
return JSONResponse({"error": "unknown-client"}, status_code=412)
|
|
@@ -970,7 +1006,7 @@ class Statewire(PinnedAPI):
|
|
|
970
1006
|
@route.websocket("/ws")
|
|
971
1007
|
async def ws(self, websocket: WebSocket) -> None:
|
|
972
1008
|
try:
|
|
973
|
-
await self.authorize(websocket)
|
|
1009
|
+
identity = await self.authorize(websocket)
|
|
974
1010
|
except HTTPException as exc:
|
|
975
1011
|
reason = exc.detail if isinstance(exc.detail, str) else "unauthorized"
|
|
976
1012
|
raise WebSocketException(code=1008, reason=reason) from None
|
|
@@ -982,7 +1018,7 @@ class Statewire(PinnedAPI):
|
|
|
982
1018
|
await websocket.accept(
|
|
983
1019
|
subprotocol=ops.WS_SUBPROTOCOL if ops.WS_SUBPROTOCOL in offered else None
|
|
984
1020
|
)
|
|
985
|
-
q, snapshot = self._register(client_id)
|
|
1021
|
+
q, snapshot = self._register(client_id, identity)
|
|
986
1022
|
lease = self._clients[client_id].lease
|
|
987
1023
|
try:
|
|
988
1024
|
await self._ws_pump(websocket, q, client_id, snapshot, lease)
|
|
@@ -20,9 +20,11 @@ the command whole):
|
|
|
20
20
|
- ``add-message`` becomes ``run/steer`` with ``{"message": <wire message>}``:
|
|
21
21
|
a fresh ``legacy_`` id is synthesized, image parts become
|
|
22
22
|
``{"type": "file", "mediaType": "image/*", "url"}`` parts, and command-level
|
|
23
|
-
anchors are dropped.
|
|
24
|
-
``
|
|
25
|
-
``
|
|
23
|
+
anchors are dropped. A command-level non-null ``sourceId`` makes it
|
|
24
|
+
``run/edit`` with that ``sourceId`` directly; otherwise, when the body
|
|
25
|
+
carries a top-level ``parentId`` and ``get_message_child_id`` resolves it to
|
|
26
|
+
a child, the command becomes ``run/edit`` with ``{"sourceId": <child>,
|
|
27
|
+
"message": <wire message>}``.
|
|
26
28
|
- ``add-tool-result`` becomes ``run/input`` with ``{"requestId": <resolved>,
|
|
27
29
|
"response": {"output": <result>, "isError"?, "artifact"?}}``; the requestId
|
|
28
30
|
comes from ``get_input_request_id(toolCallId)``, whose default peeks the
|
|
@@ -33,16 +35,19 @@ the command whole):
|
|
|
33
35
|
(None rejects).
|
|
34
36
|
|
|
35
37
|
A host without the target ``run/*`` handler rejects the batch via the
|
|
36
|
-
ordinary unknown-command path. The response streams
|
|
37
|
-
|
|
38
|
-
|
|
39
|
-
|
|
40
|
-
|
|
41
|
-
|
|
38
|
+
ordinary unknown-command path. The response streams state changes as legacy
|
|
39
|
+
frames until the instance is idle — every submitted command settled and no
|
|
40
|
+
live instance task (``create_task`` minus ``unref``) — then EOF. A run error
|
|
41
|
+
recorded into state (``get_run_error``; default ``state["error"]``) becomes
|
|
42
|
+
an error frame before EOF, its fields beyond ``message`` carried as the
|
|
43
|
+
frame's structured payload where the format allows; a rejection or crash
|
|
44
|
+
becomes an error frame followed by EOF. Client disconnect before EOF invokes
|
|
45
|
+
``on_assistant_transport_disconnect`` (a plain overridable hook, not a
|
|
46
|
+
command; default no-op); the run itself keeps executing server-side.
|
|
42
47
|
|
|
43
48
|
``POST /assistant-transport/api/resume`` (same body shape, ``commands`` must
|
|
44
49
|
be absent or empty) reattaches to the run: a full-state ``set`` at the root,
|
|
45
|
-
then the live tail until the
|
|
50
|
+
then the live tail until the instance is idle. If no run ever started, the
|
|
46
51
|
response is 200 with an empty body and ``X-Stream-Status: not_found``. If the
|
|
47
52
|
run already completed, the replay carries ``X-Stream-Status: completed``. A
|
|
48
53
|
resume disconnect never invokes the cancel hook.
|
|
@@ -56,7 +61,7 @@ existed at request time.
|
|
|
56
61
|
``POST /assistant-transport/api/status`` responds 200 always:
|
|
57
62
|
``{"isRunning": true, "status": "running"}`` while a run is active,
|
|
58
63
|
``{"isRunning": false, "status": "completed", "completedAt": <ms epoch>}``
|
|
59
|
-
|
|
64
|
+
once idle, and ``{"isRunning": false, "status": "not_found",
|
|
60
65
|
"message": ...}`` when no run ever started on this instance.
|
|
61
66
|
|
|
62
67
|
Two stream formats, selected by ``assistant_transport_protocol``:
|
|
@@ -74,7 +79,7 @@ import contextlib
|
|
|
74
79
|
import json
|
|
75
80
|
import secrets
|
|
76
81
|
import time
|
|
77
|
-
from typing import Any, AsyncIterator
|
|
82
|
+
from typing import Any, AsyncIterator, Coroutine
|
|
78
83
|
from uuid import uuid4
|
|
79
84
|
|
|
80
85
|
import httpx
|
|
@@ -102,7 +107,7 @@ class _DataStreamFormat:
|
|
|
102
107
|
return f"aui-state:[{','.join(op_frames)}]\n"
|
|
103
108
|
|
|
104
109
|
@staticmethod
|
|
105
|
-
def error(message: str) -> str:
|
|
110
|
+
def error(message: str, payload: Any = None) -> str:
|
|
106
111
|
return f"3:{json.dumps(message)}\n"
|
|
107
112
|
|
|
108
113
|
|
|
@@ -119,9 +124,11 @@ class _AssistantTransportFormat:
|
|
|
119
124
|
)
|
|
120
125
|
|
|
121
126
|
@staticmethod
|
|
122
|
-
def error(message: str) -> str:
|
|
123
|
-
|
|
124
|
-
|
|
127
|
+
def error(message: str, payload: Any = None) -> str:
|
|
128
|
+
chunk: dict[str, Any] = {"type": "error", "error": message}
|
|
129
|
+
if payload is not None:
|
|
130
|
+
chunk["payload"] = payload
|
|
131
|
+
return f"data: {json.dumps(chunk)}\n\n"
|
|
125
132
|
|
|
126
133
|
|
|
127
134
|
_FORMATS = {
|
|
@@ -183,7 +190,9 @@ _AT_CONSUMED_FIELDS = {"commands", "threadId", "state"}
|
|
|
183
190
|
|
|
184
191
|
|
|
185
192
|
class _RunError(Exception):
|
|
186
|
-
|
|
193
|
+
def __init__(self, message: str, *, payload: Any = None) -> None:
|
|
194
|
+
super().__init__(message)
|
|
195
|
+
self.payload = payload
|
|
187
196
|
|
|
188
197
|
|
|
189
198
|
def _legacy_wire_message(message: Any) -> dict[str, Any]:
|
|
@@ -215,6 +224,7 @@ class _Run:
|
|
|
215
224
|
def __init__(self, replica: Any) -> None:
|
|
216
225
|
self.replica = replica
|
|
217
226
|
self.error: str | None = None
|
|
227
|
+
self.error_payload: Any = None
|
|
218
228
|
self.done = False
|
|
219
229
|
self.completed_at: int | None = None
|
|
220
230
|
self.attachments: list["asyncio.Queue[_Item]"] = []
|
|
@@ -227,7 +237,7 @@ class _Run:
|
|
|
227
237
|
self.attachments.remove(aq)
|
|
228
238
|
while not aq.empty():
|
|
229
239
|
aq.get_nowait()
|
|
230
|
-
aq.put_nowait(("error", "state stream overflowed"))
|
|
240
|
+
aq.put_nowait(("error", ("state stream overflowed", None)))
|
|
231
241
|
aq.put_nowait(("end", None))
|
|
232
242
|
|
|
233
243
|
|
|
@@ -236,6 +246,34 @@ class AssistantTransport(Statewire):
|
|
|
236
246
|
|
|
237
247
|
_at_run: _Run | None = None
|
|
238
248
|
|
|
249
|
+
def __init__(self, ctx: Any) -> None:
|
|
250
|
+
super().__init__(ctx)
|
|
251
|
+
self._at_tasks: set["asyncio.Task[Any]"] = set()
|
|
252
|
+
|
|
253
|
+
def create_task(self, coro: Coroutine[Any, Any, Any]) -> "asyncio.Task[Any]":
|
|
254
|
+
task = super().create_task(coro)
|
|
255
|
+
self._at_tasks.add(task)
|
|
256
|
+
task.add_done_callback(self._at_tasks.discard)
|
|
257
|
+
return task
|
|
258
|
+
|
|
259
|
+
def unref(self, task: "asyncio.Task[Any]") -> "asyncio.Task[Any]":
|
|
260
|
+
self._at_tasks.discard(task)
|
|
261
|
+
return super().unref(task)
|
|
262
|
+
|
|
263
|
+
async def get_run_error(self) -> dict[str, Any] | None:
|
|
264
|
+
"""The run error to surface as the end-of-stream error frame.
|
|
265
|
+
|
|
266
|
+
The default peeks the RunManager root-state mount: a non-null
|
|
267
|
+
``state["error"]`` object with a string ``message`` is the recorded
|
|
268
|
+
run error; its fields beyond ``message`` travel as the frame's
|
|
269
|
+
structured payload where the format carries one. None ends the
|
|
270
|
+
stream cleanly."""
|
|
271
|
+
state = plain(self.state)
|
|
272
|
+
error = state.get("error") if isinstance(state, dict) else None
|
|
273
|
+
if isinstance(error, dict) and isinstance(error.get("message"), str):
|
|
274
|
+
return error
|
|
275
|
+
return None
|
|
276
|
+
|
|
239
277
|
async def on_assistant_transport_disconnect(self) -> None:
|
|
240
278
|
"""The legacy client went away or asked to stop; treat as a cancel signal."""
|
|
241
279
|
|
|
@@ -270,7 +308,7 @@ class AssistantTransport(Statewire):
|
|
|
270
308
|
|
|
271
309
|
@route.post("/assistant-transport/api/chat")
|
|
272
310
|
async def assistant_transport_chat(self, request: Request) -> Response:
|
|
273
|
-
await self.authorize(request)
|
|
311
|
+
identity = await self.authorize(request)
|
|
274
312
|
fmt = _FORMATS[self.assistant_transport_protocol]
|
|
275
313
|
raw = await request.body()
|
|
276
314
|
try:
|
|
@@ -299,7 +337,7 @@ class AssistantTransport(Statewire):
|
|
|
299
337
|
status_code=400, detail=f"invalid context: {exc}"
|
|
300
338
|
) from None
|
|
301
339
|
client_id = "at-" + secrets.token_urlsafe(9)
|
|
302
|
-
q, snapshot = self._register(client_id)
|
|
340
|
+
q, snapshot = self._register(client_id, identity)
|
|
303
341
|
self._clients[client_id].context = context
|
|
304
342
|
lease = self._clients[client_id].lease
|
|
305
343
|
assert lease is not None
|
|
@@ -308,11 +346,12 @@ class AssistantTransport(Statewire):
|
|
|
308
346
|
parent_id = None
|
|
309
347
|
run = _Run(json.loads(snapshot)["ops"][0]["value"])
|
|
310
348
|
self._at_run = run
|
|
311
|
-
self.create_task(
|
|
349
|
+
watcher = self.create_task(
|
|
312
350
|
self._assistant_transport_run(
|
|
313
351
|
run, q, client_id, lease, commands, parent_id
|
|
314
352
|
)
|
|
315
353
|
)
|
|
354
|
+
self._at_tasks.discard(watcher)
|
|
316
355
|
return StreamingResponse(
|
|
317
356
|
self._assistant_transport_attach(fmt, run, initial=True),
|
|
318
357
|
media_type="text/event-stream",
|
|
@@ -352,13 +391,13 @@ class AssistantTransport(Statewire):
|
|
|
352
391
|
|
|
353
392
|
@route.post("/assistant-transport/api/cancel")
|
|
354
393
|
async def assistant_transport_cancel(self, request: Request) -> dict[str, Any]:
|
|
355
|
-
await self.authorize(request)
|
|
394
|
+
identity = await self.authorize(request)
|
|
356
395
|
run = self._at_run
|
|
357
396
|
found = run is not None and not run.done
|
|
358
397
|
if found:
|
|
359
398
|
await self.on_assistant_transport_disconnect()
|
|
360
399
|
if "run/stop" in type(self)._statewire_commands:
|
|
361
|
-
await self._assistant_transport_dispatch("run/stop")
|
|
400
|
+
await self._assistant_transport_dispatch("run/stop", identity)
|
|
362
401
|
return {"success": True, "found": found}
|
|
363
402
|
|
|
364
403
|
@route.post("/assistant-transport/api/status")
|
|
@@ -415,11 +454,11 @@ class AssistantTransport(Statewire):
|
|
|
415
454
|
registered = type(self)._statewire_commands
|
|
416
455
|
if kind == "add-message" and "add-message" not in registered:
|
|
417
456
|
message = _legacy_wire_message(command.get("message"))
|
|
418
|
-
source_id = (
|
|
419
|
-
|
|
420
|
-
|
|
421
|
-
|
|
422
|
-
|
|
457
|
+
source_id = command.get("sourceId")
|
|
458
|
+
if source_id is not None and not isinstance(source_id, str):
|
|
459
|
+
raise _RunError("add-message: sourceId must be a string or null")
|
|
460
|
+
if source_id is None and parent_id is not None:
|
|
461
|
+
source_id = await self.get_message_child_id(parent_id)
|
|
423
462
|
if source_id is not None:
|
|
424
463
|
return {
|
|
425
464
|
"method": "run/edit",
|
|
@@ -457,9 +496,9 @@ class AssistantTransport(Statewire):
|
|
|
457
496
|
raise _RunError(f"reload: no child message for parentId {parent_id!r}")
|
|
458
497
|
return {"method": "run/reload", "params": [{"sourceId": source_id}]}
|
|
459
498
|
|
|
460
|
-
async def _assistant_transport_dispatch(self, method: str) -> None:
|
|
499
|
+
async def _assistant_transport_dispatch(self, method: str, identity: Any) -> None:
|
|
461
500
|
client_id = "at-" + secrets.token_urlsafe(9)
|
|
462
|
-
q, _ = self._register(client_id)
|
|
501
|
+
q, _ = self._register(client_id, identity)
|
|
463
502
|
try:
|
|
464
503
|
lease = self._clients[client_id].lease
|
|
465
504
|
assert lease is not None
|
|
@@ -486,48 +525,82 @@ class AssistantTransport(Statewire):
|
|
|
486
525
|
total = len(commands) or (1 if parent_id is not None else 0)
|
|
487
526
|
acked = 0
|
|
488
527
|
open_pending: set[int] = set()
|
|
528
|
+
|
|
529
|
+
def idle() -> bool:
|
|
530
|
+
return (
|
|
531
|
+
submit.done()
|
|
532
|
+
and acked >= total
|
|
533
|
+
and not open_pending
|
|
534
|
+
and not self._at_tasks
|
|
535
|
+
)
|
|
536
|
+
|
|
537
|
+
def consume(item: Any) -> bool:
|
|
538
|
+
nonlocal acked
|
|
539
|
+
if item is _LAGGARD:
|
|
540
|
+
raise _RunError("state stream overflowed")
|
|
541
|
+
data, finish = item
|
|
542
|
+
envelope = json.loads(data)
|
|
543
|
+
op_frames: list[str] = []
|
|
544
|
+
for op in envelope.get("ops", []):
|
|
545
|
+
run.replica, frame = _translate_op(run.replica, op)
|
|
546
|
+
if frame is not None:
|
|
547
|
+
op_frames.append(frame)
|
|
548
|
+
if op_frames:
|
|
549
|
+
run.push(("state", op_frames))
|
|
550
|
+
for res in envelope.get("res", []):
|
|
551
|
+
if res["type"] == "pending":
|
|
552
|
+
open_pending.add(res["seq"])
|
|
553
|
+
continue
|
|
554
|
+
open_pending.discard(res["seq"])
|
|
555
|
+
if res["type"] in ("rejected", "crashed"):
|
|
556
|
+
raise _RunError(
|
|
557
|
+
res.get("message") or res["type"],
|
|
558
|
+
payload=res.get("payload"),
|
|
559
|
+
)
|
|
560
|
+
if "ack" in envelope:
|
|
561
|
+
acked = envelope["ack"]
|
|
562
|
+
if finish and not idle():
|
|
563
|
+
fin = envelope.get("fin", {})
|
|
564
|
+
raise _RunError(
|
|
565
|
+
fin.get("message") or fin.get("reason", "stream ended")
|
|
566
|
+
)
|
|
567
|
+
return finish
|
|
568
|
+
|
|
489
569
|
try:
|
|
490
|
-
|
|
491
|
-
|
|
570
|
+
ended = False
|
|
571
|
+
while not ended:
|
|
572
|
+
if idle():
|
|
573
|
+
self.drain()
|
|
574
|
+
self.flush()
|
|
575
|
+
if q_get.done():
|
|
576
|
+
ended = consume(q_get.result()) or ended
|
|
577
|
+
q_get = asyncio.create_task(q.get())
|
|
578
|
+
continue
|
|
579
|
+
while not q.empty():
|
|
580
|
+
ended = consume(q.get_nowait()) or ended
|
|
581
|
+
if idle():
|
|
582
|
+
break
|
|
583
|
+
continue
|
|
584
|
+
waiters = {q_get, *self._at_tasks}
|
|
585
|
+
if not submit.done():
|
|
586
|
+
waiters.add(submit)
|
|
492
587
|
done, _ = await asyncio.wait(
|
|
493
588
|
waiters, return_when=asyncio.FIRST_COMPLETED
|
|
494
589
|
)
|
|
495
590
|
if submit in done and submit.exception() is not None:
|
|
496
591
|
raise _RunError(str(submit.exception()))
|
|
497
|
-
if q_get
|
|
498
|
-
|
|
499
|
-
|
|
500
|
-
|
|
501
|
-
|
|
502
|
-
|
|
503
|
-
|
|
504
|
-
|
|
505
|
-
for op in envelope.get("ops", []):
|
|
506
|
-
run.replica, frame = _translate_op(run.replica, op)
|
|
507
|
-
if frame is not None:
|
|
508
|
-
op_frames.append(frame)
|
|
509
|
-
if op_frames:
|
|
510
|
-
run.push(("state", op_frames))
|
|
511
|
-
for res in envelope.get("res", []):
|
|
512
|
-
if res["type"] == "pending":
|
|
513
|
-
open_pending.add(res["seq"])
|
|
514
|
-
continue
|
|
515
|
-
open_pending.discard(res["seq"])
|
|
516
|
-
if res["type"] in ("rejected", "crashed"):
|
|
517
|
-
raise _RunError(res.get("message") or res["type"])
|
|
518
|
-
if "ack" in envelope:
|
|
519
|
-
acked = envelope["ack"]
|
|
520
|
-
if finish:
|
|
521
|
-
if submit.done() and acked >= total and not open_pending:
|
|
522
|
-
break
|
|
523
|
-
fin = envelope.get("fin", {})
|
|
524
|
-
raise _RunError(
|
|
525
|
-
fin.get("message") or fin.get("reason", "stream ended")
|
|
526
|
-
)
|
|
527
|
-
q_get = asyncio.create_task(q.get())
|
|
592
|
+
if q_get in done:
|
|
593
|
+
ended = consume(q_get.result())
|
|
594
|
+
if not ended:
|
|
595
|
+
q_get = asyncio.create_task(q.get())
|
|
596
|
+
error = await self.get_run_error()
|
|
597
|
+
if error is not None:
|
|
598
|
+
payload = {k: v for k, v in error.items() if k != "message"}
|
|
599
|
+
raise _RunError(error["message"], payload=payload or None)
|
|
528
600
|
except _RunError as exc:
|
|
529
601
|
run.error = str(exc)
|
|
530
|
-
run.
|
|
602
|
+
run.error_payload = exc.payload
|
|
603
|
+
run.push(("error", (run.error, run.error_payload)))
|
|
531
604
|
finally:
|
|
532
605
|
q_get.cancel()
|
|
533
606
|
submit.cancel()
|
|
@@ -559,11 +632,11 @@ class AssistantTransport(Statewire):
|
|
|
559
632
|
if kind == "state":
|
|
560
633
|
yield fmt.state(value)
|
|
561
634
|
elif kind == "error":
|
|
562
|
-
yield fmt.error(value)
|
|
635
|
+
yield fmt.error(*value)
|
|
563
636
|
else:
|
|
564
637
|
break
|
|
565
638
|
elif initial and run.error is not None:
|
|
566
|
-
yield fmt.error(run.error)
|
|
639
|
+
yield fmt.error(run.error, run.error_payload)
|
|
567
640
|
if fmt.end:
|
|
568
641
|
yield fmt.end
|
|
569
642
|
completed = True
|
|
@@ -2,13 +2,16 @@
|
|
|
2
2
|
|
|
3
3
|
The legacy client POSTs ``{"commands", "state", "threadId", ...}`` and reads
|
|
4
4
|
back a data-stream (``<prefix>:<json>`` lines) or assistant-transport SSE
|
|
5
|
-
response carrying ``set``/``append-text`` state ops until the
|
|
5
|
+
response carrying ``set``/``append-text`` state ops until the instance is
|
|
6
|
+
idle.
|
|
6
7
|
"""
|
|
7
8
|
|
|
8
9
|
import asyncio
|
|
9
10
|
import contextlib
|
|
10
11
|
import json
|
|
11
12
|
|
|
13
|
+
import pytest
|
|
14
|
+
|
|
12
15
|
from statewire_helpers import statewire_client
|
|
13
16
|
|
|
14
17
|
from statewire import AssistantTransport, StatewireReject, command
|
|
@@ -110,7 +113,7 @@ async def post_run(client, body, thread="t1"):
|
|
|
110
113
|
return await client.post(f"/threads/{thread}/assistant-transport/api/chat", json=body)
|
|
111
114
|
|
|
112
115
|
|
|
113
|
-
async def
|
|
116
|
+
async def test_data_stream_response_streams_snapshot_then_ops_until_idle():
|
|
114
117
|
async with statewire_client(Harness) as (app, client):
|
|
115
118
|
response = await post_run(
|
|
116
119
|
client,
|
|
@@ -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
|
|
@@ -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
|
|
|
@@ -992,6 +1135,58 @@ async def test_add_message_with_a_resolvable_parent_is_normalized_to_run_edit():
|
|
|
992
1135
|
assert params["message"]["parts"] == [{"type": "text", "text": "again"}]
|
|
993
1136
|
|
|
994
1137
|
|
|
1138
|
+
async def test_add_message_with_a_source_id_is_normalized_to_run_edit_with_it():
|
|
1139
|
+
async with statewire_client(AnchoredHarness) as (app, client):
|
|
1140
|
+
response = await post_run(
|
|
1141
|
+
client,
|
|
1142
|
+
{
|
|
1143
|
+
"commands": [
|
|
1144
|
+
{
|
|
1145
|
+
"type": "add-message",
|
|
1146
|
+
"message": {
|
|
1147
|
+
"role": "user",
|
|
1148
|
+
"parts": [{"type": "text", "text": "again"}],
|
|
1149
|
+
},
|
|
1150
|
+
"sourceId": "m9",
|
|
1151
|
+
}
|
|
1152
|
+
],
|
|
1153
|
+
"state": {},
|
|
1154
|
+
"parentId": "m1",
|
|
1155
|
+
},
|
|
1156
|
+
)
|
|
1157
|
+
assert response.status_code == 200
|
|
1158
|
+
instance = (await app.state.pinned_host.directory.get("t1")).instance
|
|
1159
|
+
assert instance.steered == []
|
|
1160
|
+
[params] = instance.edited
|
|
1161
|
+
assert params["sourceId"] == "m9"
|
|
1162
|
+
assert params["message"]["parts"] == [{"type": "text", "text": "again"}]
|
|
1163
|
+
|
|
1164
|
+
|
|
1165
|
+
async def test_add_message_with_a_null_source_id_and_no_parent_is_run_steer():
|
|
1166
|
+
async with statewire_client(AnchoredHarness) as (app, client):
|
|
1167
|
+
response = await post_run(
|
|
1168
|
+
client,
|
|
1169
|
+
{
|
|
1170
|
+
"commands": [
|
|
1171
|
+
{
|
|
1172
|
+
"type": "add-message",
|
|
1173
|
+
"message": {
|
|
1174
|
+
"role": "user",
|
|
1175
|
+
"parts": [{"type": "text", "text": "hi"}],
|
|
1176
|
+
},
|
|
1177
|
+
"sourceId": None,
|
|
1178
|
+
}
|
|
1179
|
+
],
|
|
1180
|
+
"state": {},
|
|
1181
|
+
},
|
|
1182
|
+
)
|
|
1183
|
+
assert response.status_code == 200
|
|
1184
|
+
instance = (await app.state.pinned_host.directory.get("t1")).instance
|
|
1185
|
+
assert instance.edited == []
|
|
1186
|
+
[(params, _)] = instance.steered
|
|
1187
|
+
assert set(params) == {"message"}
|
|
1188
|
+
|
|
1189
|
+
|
|
995
1190
|
async def test_add_message_with_an_unresolvable_parent_falls_back_to_run_steer():
|
|
996
1191
|
async with statewire_client(AnchoredHarness) as (app, client):
|
|
997
1192
|
response = await post_run(
|
|
@@ -0,0 +1,256 @@
|
|
|
1
|
+
from statewire_helpers import post_command, statewire_client, stream_of, ws_of
|
|
2
|
+
from fastapi import HTTPException, Request
|
|
3
|
+
|
|
4
|
+
from statewire import AssistantTransport, Statewire, command
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
instances = []
|
|
8
|
+
runs = []
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
class Guarded(Statewire):
|
|
12
|
+
async def lifespan(self):
|
|
13
|
+
instances.append(self)
|
|
14
|
+
self.state = {"count": 0}
|
|
15
|
+
yield
|
|
16
|
+
|
|
17
|
+
async def authorize(self, request: Request) -> None:
|
|
18
|
+
if request.headers.get("x-token") != "secret":
|
|
19
|
+
raise HTTPException(status_code=401, detail="unauthorized")
|
|
20
|
+
|
|
21
|
+
@command
|
|
22
|
+
async def inc(self):
|
|
23
|
+
runs.append("inc")
|
|
24
|
+
self.state["count"] += 1
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
async def test_stream_attach_without_token_is_401_and_registers_nothing():
|
|
28
|
+
async with statewire_client(Guarded) as (app, client):
|
|
29
|
+
async with stream_of(app) as stream:
|
|
30
|
+
assert stream.status == 401
|
|
31
|
+
assert instances[-1]._subscribers == {}
|
|
32
|
+
|
|
33
|
+
|
|
34
|
+
async def test_command_without_token_is_401_and_handler_never_runs():
|
|
35
|
+
async with statewire_client(Guarded) as (app, client):
|
|
36
|
+
async with stream_of(app, headers={"x-token": "secret"}) as stream:
|
|
37
|
+
await stream.next_event()
|
|
38
|
+
before = len(runs)
|
|
39
|
+
response = await post_command(client, {"method": "inc", "params": []}, seq=1)
|
|
40
|
+
assert response.status_code == 401
|
|
41
|
+
assert len(runs) == before
|
|
42
|
+
|
|
43
|
+
|
|
44
|
+
async def test_authorized_requests_succeed():
|
|
45
|
+
async with statewire_client(Guarded) as (app, client):
|
|
46
|
+
async with stream_of(app, headers={"x-token": "secret"}) as stream:
|
|
47
|
+
assert stream.status == 200
|
|
48
|
+
await stream.next_event()
|
|
49
|
+
response = await post_command(
|
|
50
|
+
client, {"method": "inc", "params": []}, seq=1, headers={"x-token": "secret"}
|
|
51
|
+
)
|
|
52
|
+
assert response.status_code == 200
|
|
53
|
+
assert await stream.next_event() == {
|
|
54
|
+
"ops": [{"op": "replace", "path": ["count"], "value": 1}],
|
|
55
|
+
"ack": 1,
|
|
56
|
+
}
|
|
57
|
+
|
|
58
|
+
|
|
59
|
+
class Identified(Statewire):
|
|
60
|
+
caller = None
|
|
61
|
+
|
|
62
|
+
async def lifespan(self):
|
|
63
|
+
self.state = {}
|
|
64
|
+
yield
|
|
65
|
+
|
|
66
|
+
async def authorize(self, request):
|
|
67
|
+
return request.headers.get("x-user")
|
|
68
|
+
|
|
69
|
+
@command
|
|
70
|
+
async def whoami(self, *, ctx):
|
|
71
|
+
self.caller = ctx.caller
|
|
72
|
+
return {"identity": ctx.caller.identity}
|
|
73
|
+
|
|
74
|
+
|
|
75
|
+
async def instance_of(app, thread="t1"):
|
|
76
|
+
return (await app.state.pinned_host.directory.get(thread)).instance
|
|
77
|
+
|
|
78
|
+
|
|
79
|
+
async def whoami_res(client, *, seq, headers=None):
|
|
80
|
+
body = {"method": "whoami", "params": []}
|
|
81
|
+
first = await post_command(client, body, seq=seq, headers=headers)
|
|
82
|
+
assert first.status_code == 200
|
|
83
|
+
dup = await post_command(client, body, seq=seq, headers=headers)
|
|
84
|
+
return dup.json()["res"]
|
|
85
|
+
|
|
86
|
+
|
|
87
|
+
async def test_identity_returned_by_authorize_is_visible_via_ctx_caller():
|
|
88
|
+
async with statewire_client(Identified) as (app, client):
|
|
89
|
+
async with stream_of(app, headers={"x-user": "alice"}) as stream:
|
|
90
|
+
await stream.next_event()
|
|
91
|
+
res = await whoami_res(client, seq=1, headers={"x-user": "alice"})
|
|
92
|
+
assert res["payload"] == {"identity": "alice"}
|
|
93
|
+
|
|
94
|
+
|
|
95
|
+
async def test_identity_is_visible_over_ws():
|
|
96
|
+
async with statewire_client(Identified) as (app, client):
|
|
97
|
+
async with ws_of(app, headers={"x-user": "wanda"}) as ws:
|
|
98
|
+
await ws.next_frame()
|
|
99
|
+
env = await ws.command({"method": "whoami", "params": []})
|
|
100
|
+
assert env["res"] == [
|
|
101
|
+
{"seq": 1, "type": "accepted", "payload": {"identity": "wanda"}}
|
|
102
|
+
]
|
|
103
|
+
|
|
104
|
+
|
|
105
|
+
async def test_reattach_restamps_the_identity_on_a_held_handle():
|
|
106
|
+
async with statewire_client(Identified) as (app, client):
|
|
107
|
+
async with stream_of(app, headers={"x-user": "alice"}) as stream:
|
|
108
|
+
await stream.next_event()
|
|
109
|
+
res = await whoami_res(client, seq=1, headers={"x-user": "alice"})
|
|
110
|
+
assert res["payload"] == {"identity": "alice"}
|
|
111
|
+
instance = await instance_of(app)
|
|
112
|
+
assert instance.caller.identity == "alice"
|
|
113
|
+
async with stream_of(app, headers={"x-user": "bob"}) as stream:
|
|
114
|
+
await stream.next_event()
|
|
115
|
+
assert instance.caller.identity == "bob"
|
|
116
|
+
|
|
117
|
+
|
|
118
|
+
async def test_identity_defaults_to_none_when_authorize_is_not_overridden():
|
|
119
|
+
class Anonymous(Statewire):
|
|
120
|
+
async def lifespan(self):
|
|
121
|
+
self.state = {}
|
|
122
|
+
yield
|
|
123
|
+
|
|
124
|
+
@command
|
|
125
|
+
async def whoami(self, *, ctx):
|
|
126
|
+
return {"identity": ctx.caller.identity}
|
|
127
|
+
|
|
128
|
+
async with statewire_client(Anonymous) as (app, client):
|
|
129
|
+
async with stream_of(app) as stream:
|
|
130
|
+
await stream.next_event()
|
|
131
|
+
body = {"method": "whoami", "params": []}
|
|
132
|
+
assert (await post_command(client, body, seq=1)).status_code == 200
|
|
133
|
+
dup = await post_command(client, body, seq=1)
|
|
134
|
+
assert dup.json()["res"]["payload"] == {"identity": None}
|
|
135
|
+
|
|
136
|
+
|
|
137
|
+
async def test_forged_client_id_with_bad_lease_cannot_overwrite_identity():
|
|
138
|
+
async with statewire_client(Identified) as (app, client):
|
|
139
|
+
async with stream_of(app, headers={"x-user": "alice"}) as stream:
|
|
140
|
+
await stream.next_event()
|
|
141
|
+
res = await whoami_res(client, seq=1, headers={"x-user": "alice"})
|
|
142
|
+
assert res["payload"] == {"identity": "alice"}
|
|
143
|
+
instance = await instance_of(app)
|
|
144
|
+
for lease in (None, "stolen"):
|
|
145
|
+
response = await post_command(
|
|
146
|
+
client,
|
|
147
|
+
{"method": "whoami", "params": []},
|
|
148
|
+
seq=2,
|
|
149
|
+
headers={"x-user": "mallory"},
|
|
150
|
+
lease=lease,
|
|
151
|
+
)
|
|
152
|
+
assert response.status_code == 423
|
|
153
|
+
assert instance.caller.identity == "alice"
|
|
154
|
+
|
|
155
|
+
|
|
156
|
+
async def test_unadmitted_lease_valid_submissions_do_not_restamp_identity():
|
|
157
|
+
async with statewire_client(Identified) as (app, client):
|
|
158
|
+
async with stream_of(app, headers={"x-user": "alice"}) as stream:
|
|
159
|
+
await stream.next_event()
|
|
160
|
+
res = await whoami_res(client, seq=1, headers={"x-user": "alice"})
|
|
161
|
+
assert res["payload"] == {"identity": "alice"}
|
|
162
|
+
instance = await instance_of(app)
|
|
163
|
+
duplicate = await post_command(
|
|
164
|
+
client,
|
|
165
|
+
{"method": "whoami", "params": []},
|
|
166
|
+
seq=1,
|
|
167
|
+
headers={"x-user": "mallory"},
|
|
168
|
+
)
|
|
169
|
+
assert duplicate.status_code == 200
|
|
170
|
+
assert instance.caller.identity == "alice"
|
|
171
|
+
gap = await post_command(
|
|
172
|
+
client,
|
|
173
|
+
{"method": "whoami", "params": []},
|
|
174
|
+
seq=9,
|
|
175
|
+
headers={"x-user": "mallory"},
|
|
176
|
+
)
|
|
177
|
+
assert gap.status_code == 409
|
|
178
|
+
assert instance.caller.identity == "alice"
|
|
179
|
+
rejected = await post_command(
|
|
180
|
+
client,
|
|
181
|
+
{"method": "no-such-method", "params": []},
|
|
182
|
+
seq=2,
|
|
183
|
+
headers={"x-user": "mallory"},
|
|
184
|
+
)
|
|
185
|
+
assert rejected.status_code == 422
|
|
186
|
+
assert instance.caller.identity == "alice"
|
|
187
|
+
|
|
188
|
+
|
|
189
|
+
async def test_identity_is_not_settable_via_context_updates():
|
|
190
|
+
async with statewire_client(Identified) as (app, client):
|
|
191
|
+
async with stream_of(app, headers={"x-user": "alice"}) as stream:
|
|
192
|
+
await stream.next_event()
|
|
193
|
+
response = await post_command(
|
|
194
|
+
client,
|
|
195
|
+
{"context": [{"op": "replace", "path": [], "value": {"identity": "hacker"}}]},
|
|
196
|
+
seq=1,
|
|
197
|
+
headers={"x-user": "alice"},
|
|
198
|
+
)
|
|
199
|
+
assert response.status_code == 200
|
|
200
|
+
res = await whoami_res(client, seq=2, headers={"x-user": "alice"})
|
|
201
|
+
assert res["payload"] == {"identity": "alice"}
|
|
202
|
+
|
|
203
|
+
|
|
204
|
+
class LegacyIdentified(AssistantTransport):
|
|
205
|
+
seen = None
|
|
206
|
+
stopped = None
|
|
207
|
+
|
|
208
|
+
async def lifespan(self):
|
|
209
|
+
self.state = {}
|
|
210
|
+
yield
|
|
211
|
+
|
|
212
|
+
async def authorize(self, request):
|
|
213
|
+
return request.headers.get("x-user")
|
|
214
|
+
|
|
215
|
+
@command("run/steer")
|
|
216
|
+
async def steer(self, params, *, ctx):
|
|
217
|
+
self.seen = ctx.caller.identity
|
|
218
|
+
|
|
219
|
+
@command("run/stop")
|
|
220
|
+
async def stop(self, *, ctx):
|
|
221
|
+
self.stopped = ctx.caller.identity
|
|
222
|
+
|
|
223
|
+
|
|
224
|
+
async def test_shim_run_client_carries_the_identity_into_translated_commands():
|
|
225
|
+
async with statewire_client(LegacyIdentified) as (app, client):
|
|
226
|
+
response = await client.post(
|
|
227
|
+
"/threads/t1/assistant-transport/api/chat",
|
|
228
|
+
headers={"x-user": "legacy"},
|
|
229
|
+
json={
|
|
230
|
+
"threadId": "t1",
|
|
231
|
+
"commands": [
|
|
232
|
+
{
|
|
233
|
+
"type": "add-message",
|
|
234
|
+
"message": {
|
|
235
|
+
"role": "user",
|
|
236
|
+
"parts": [{"type": "text", "text": "hi"}],
|
|
237
|
+
},
|
|
238
|
+
}
|
|
239
|
+
],
|
|
240
|
+
},
|
|
241
|
+
)
|
|
242
|
+
assert response.status_code == 200
|
|
243
|
+
instance = await instance_of(app)
|
|
244
|
+
assert instance.seen == "legacy"
|
|
245
|
+
|
|
246
|
+
|
|
247
|
+
async def test_shim_cancel_dispatch_carries_the_identity():
|
|
248
|
+
async with statewire_client(LegacyIdentified) as (app, client):
|
|
249
|
+
response = await client.post(
|
|
250
|
+
"/threads/t1/assistant-transport/api/cancel",
|
|
251
|
+
headers={"x-user": "canceller"},
|
|
252
|
+
json={"threadId": "t1"},
|
|
253
|
+
)
|
|
254
|
+
assert response.status_code == 200
|
|
255
|
+
instance = await instance_of(app)
|
|
256
|
+
assert instance.stopped == "canceller"
|
|
@@ -258,7 +258,7 @@ async def test_get_value_matches_assistant_streams_unwrap_pattern():
|
|
|
258
258
|
|
|
259
259
|
async def test_flush_delivers_the_pending_envelope_synchronously():
|
|
260
260
|
inst = make({"a": 1})
|
|
261
|
-
q, _ = inst._register("c1")
|
|
261
|
+
q, _ = inst._register("c1", None)
|
|
262
262
|
inst.state["a"] = 2
|
|
263
263
|
assert q.empty()
|
|
264
264
|
inst.flush()
|
|
@@ -87,7 +87,7 @@ async def test_ws_command_on_a_superseded_socket_closes_1008_without_executing()
|
|
|
87
87
|
await ws.next_frame()
|
|
88
88
|
record = await app.state.pinned_host.directory.get("t1")
|
|
89
89
|
await ws.send_frame({"method": "inc", "params": [], "seq": 1})
|
|
90
|
-
record.instance._register("c1")
|
|
90
|
+
record.instance._register("c1", None)
|
|
91
91
|
while await ws.next_frame() is not None:
|
|
92
92
|
pass
|
|
93
93
|
assert ws.close_code == 1008
|
|
@@ -1,56 +0,0 @@
|
|
|
1
|
-
from statewire_helpers import post_command, statewire_client, stream_of
|
|
2
|
-
from fastapi import HTTPException, Request
|
|
3
|
-
|
|
4
|
-
from statewire import Statewire, command
|
|
5
|
-
|
|
6
|
-
|
|
7
|
-
instances = []
|
|
8
|
-
runs = []
|
|
9
|
-
|
|
10
|
-
|
|
11
|
-
class Guarded(Statewire):
|
|
12
|
-
async def lifespan(self):
|
|
13
|
-
instances.append(self)
|
|
14
|
-
self.state = {"count": 0}
|
|
15
|
-
yield
|
|
16
|
-
|
|
17
|
-
async def authorize(self, request: Request) -> None:
|
|
18
|
-
if request.headers.get("x-token") != "secret":
|
|
19
|
-
raise HTTPException(status_code=401, detail="unauthorized")
|
|
20
|
-
|
|
21
|
-
@command
|
|
22
|
-
async def inc(self):
|
|
23
|
-
runs.append("inc")
|
|
24
|
-
self.state["count"] += 1
|
|
25
|
-
|
|
26
|
-
|
|
27
|
-
async def test_stream_attach_without_token_is_401_and_registers_nothing():
|
|
28
|
-
async with statewire_client(Guarded) as (app, client):
|
|
29
|
-
async with stream_of(app) as stream:
|
|
30
|
-
assert stream.status == 401
|
|
31
|
-
assert instances[-1]._subscribers == {}
|
|
32
|
-
|
|
33
|
-
|
|
34
|
-
async def test_command_without_token_is_401_and_handler_never_runs():
|
|
35
|
-
async with statewire_client(Guarded) as (app, client):
|
|
36
|
-
async with stream_of(app, headers={"x-token": "secret"}) as stream:
|
|
37
|
-
await stream.next_event()
|
|
38
|
-
before = len(runs)
|
|
39
|
-
response = await post_command(client, {"method": "inc", "params": []}, seq=1)
|
|
40
|
-
assert response.status_code == 401
|
|
41
|
-
assert len(runs) == before
|
|
42
|
-
|
|
43
|
-
|
|
44
|
-
async def test_authorized_requests_succeed():
|
|
45
|
-
async with statewire_client(Guarded) as (app, client):
|
|
46
|
-
async with stream_of(app, headers={"x-token": "secret"}) as stream:
|
|
47
|
-
assert stream.status == 200
|
|
48
|
-
await stream.next_event()
|
|
49
|
-
response = await post_command(
|
|
50
|
-
client, {"method": "inc", "params": []}, seq=1, headers={"x-token": "secret"}
|
|
51
|
-
)
|
|
52
|
-
assert response.status_code == 200
|
|
53
|
-
assert await stream.next_event() == {
|
|
54
|
-
"ops": [{"op": "replace", "path": ["count"], "value": 1}],
|
|
55
|
-
"ack": 1,
|
|
56
|
-
}
|
|
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
|