statewire 0.3.1__tar.gz → 0.3.2__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.2}/PKG-INFO +1 -1
- {statewire-0.3.1 → statewire-0.3.2}/pyproject.toml +1 -1
- {statewire-0.3.1 → statewire-0.3.2}/src/statewire/api.py +50 -14
- {statewire-0.3.1 → statewire-0.3.2}/src/statewire/assistant_transport.py +6 -6
- statewire-0.3.2/tests/test_authorize.py +256 -0
- {statewire-0.3.1 → statewire-0.3.2}/tests/test_state_proxy.py +1 -1
- {statewire-0.3.1 → statewire-0.3.2}/tests/test_writer_lease.py +1 -1
- statewire-0.3.1/tests/test_authorize.py +0 -56
- {statewire-0.3.1 → statewire-0.3.2}/.gitignore +0 -0
- {statewire-0.3.1 → statewire-0.3.2}/README.md +0 -0
- {statewire-0.3.1 → statewire-0.3.2}/examples/__init__.py +0 -0
- {statewire-0.3.1 → statewire-0.3.2}/examples/demo_app.py +0 -0
- {statewire-0.3.1 → statewire-0.3.2}/src/statewire/__init__.py +0 -0
- {statewire-0.3.1 → statewire-0.3.2}/src/statewire/assistant_transport_client.py +0 -0
- {statewire-0.3.1 → statewire-0.3.2}/src/statewire/client.py +0 -0
- {statewire-0.3.1 → statewire-0.3.2}/src/statewire/langgraph.py +0 -0
- {statewire-0.3.1 → statewire-0.3.2}/src/statewire/ops.py +0 -0
- {statewire-0.3.1 → statewire-0.3.2}/src/statewire/state.py +0 -0
- {statewire-0.3.1 → statewire-0.3.2}/tests/client_helpers.py +0 -0
- {statewire-0.3.1 → statewire-0.3.2}/tests/statewire_helpers.py +0 -0
- {statewire-0.3.1 → statewire-0.3.2}/tests/test_assistant_transport.py +0 -0
- {statewire-0.3.1 → statewire-0.3.2}/tests/test_assistant_transport_client.py +0 -0
- {statewire-0.3.1 → statewire-0.3.2}/tests/test_assistant_transport_facade.py +0 -0
- {statewire-0.3.1 → statewire-0.3.2}/tests/test_client.py +0 -0
- {statewire-0.3.1 → statewire-0.3.2}/tests/test_client_ws.py +0 -0
- {statewire-0.3.1 → statewire-0.3.2}/tests/test_commands.py +0 -0
- {statewire-0.3.1 → statewire-0.3.2}/tests/test_context.py +0 -0
- {statewire-0.3.1 → statewire-0.3.2}/tests/test_langgraph.py +0 -0
- {statewire-0.3.1 → statewire-0.3.2}/tests/test_lifespan.py +0 -0
- {statewire-0.3.1 → statewire-0.3.2}/tests/test_meta.py +0 -0
- {statewire-0.3.1 → statewire-0.3.2}/tests/test_statewire_hostable.py +0 -0
- {statewire-0.3.1 → statewire-0.3.2}/tests/test_stream.py +0 -0
- {statewire-0.3.1 → statewire-0.3.2}/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)
|
|
@@ -270,7 +270,7 @@ class AssistantTransport(Statewire):
|
|
|
270
270
|
|
|
271
271
|
@route.post("/assistant-transport/api/chat")
|
|
272
272
|
async def assistant_transport_chat(self, request: Request) -> Response:
|
|
273
|
-
await self.authorize(request)
|
|
273
|
+
identity = await self.authorize(request)
|
|
274
274
|
fmt = _FORMATS[self.assistant_transport_protocol]
|
|
275
275
|
raw = await request.body()
|
|
276
276
|
try:
|
|
@@ -299,7 +299,7 @@ class AssistantTransport(Statewire):
|
|
|
299
299
|
status_code=400, detail=f"invalid context: {exc}"
|
|
300
300
|
) from None
|
|
301
301
|
client_id = "at-" + secrets.token_urlsafe(9)
|
|
302
|
-
q, snapshot = self._register(client_id)
|
|
302
|
+
q, snapshot = self._register(client_id, identity)
|
|
303
303
|
self._clients[client_id].context = context
|
|
304
304
|
lease = self._clients[client_id].lease
|
|
305
305
|
assert lease is not None
|
|
@@ -352,13 +352,13 @@ class AssistantTransport(Statewire):
|
|
|
352
352
|
|
|
353
353
|
@route.post("/assistant-transport/api/cancel")
|
|
354
354
|
async def assistant_transport_cancel(self, request: Request) -> dict[str, Any]:
|
|
355
|
-
await self.authorize(request)
|
|
355
|
+
identity = await self.authorize(request)
|
|
356
356
|
run = self._at_run
|
|
357
357
|
found = run is not None and not run.done
|
|
358
358
|
if found:
|
|
359
359
|
await self.on_assistant_transport_disconnect()
|
|
360
360
|
if "run/stop" in type(self)._statewire_commands:
|
|
361
|
-
await self._assistant_transport_dispatch("run/stop")
|
|
361
|
+
await self._assistant_transport_dispatch("run/stop", identity)
|
|
362
362
|
return {"success": True, "found": found}
|
|
363
363
|
|
|
364
364
|
@route.post("/assistant-transport/api/status")
|
|
@@ -457,9 +457,9 @@ class AssistantTransport(Statewire):
|
|
|
457
457
|
raise _RunError(f"reload: no child message for parentId {parent_id!r}")
|
|
458
458
|
return {"method": "run/reload", "params": [{"sourceId": source_id}]}
|
|
459
459
|
|
|
460
|
-
async def _assistant_transport_dispatch(self, method: str) -> None:
|
|
460
|
+
async def _assistant_transport_dispatch(self, method: str, identity: Any) -> None:
|
|
461
461
|
client_id = "at-" + secrets.token_urlsafe(9)
|
|
462
|
-
q, _ = self._register(client_id)
|
|
462
|
+
q, _ = self._register(client_id, identity)
|
|
463
463
|
try:
|
|
464
464
|
lease = self._clients[client_id].lease
|
|
465
465
|
assert lease is not None
|
|
@@ -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
|
|
File without changes
|