inter-agent-core 0.2.0__py3-none-any.whl
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.
- inter_agent/__init__.py +0 -0
- inter_agent/core/__init__.py +1 -0
- inter_agent/core/adapter_control.py +350 -0
- inter_agent/core/auth.py +279 -0
- inter_agent/core/channels.py +131 -0
- inter_agent/core/client.py +345 -0
- inter_agent/core/config.py +325 -0
- inter_agent/core/errors.py +31 -0
- inter_agent/core/kick.py +111 -0
- inter_agent/core/list.py +127 -0
- inter_agent/core/publish.py +146 -0
- inter_agent/core/router.py +8 -0
- inter_agent/core/send.py +265 -0
- inter_agent/core/server.py +765 -0
- inter_agent/core/shared.py +182 -0
- inter_agent/core/shutdown.py +92 -0
- inter_agent/core/status.py +291 -0
- inter_agent/core/tls.py +141 -0
- inter_agent/core/transport.py +37 -0
- inter_agent/py.typed +0 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/asyncapi.yaml +225 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/error-codes.md +28 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/examples/auth_challenge.json +6 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/examples/auth_response.json +4 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/examples/broadcast.json +4 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/examples/bye.json +3 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/examples/channels.json +3 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/examples/channels_ok.json +9 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/examples/custom.unknown-pass-through.json +9 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/examples/error.auth-failed.json +5 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/examples/hello.agent-with-label.json +14 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/examples/hello.agent.json +13 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/examples/kick.json +4 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/examples/kick_ok.json +5 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/examples/list.json +3 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/examples/list_ok.json +10 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/examples/msg.custom.json +13 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/examples/msg.text.json +9 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/examples/ping.json +3 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/examples/pong.json +3 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/examples/publish.json +5 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/examples/send.direct.json +5 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/examples/shutdown.json +3 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/examples/shutdown_ok.json +3 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/examples/subscribe.json +4 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/examples/subscribe_ok.json +4 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/examples/unsubscribe.json +4 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/examples/unsubscribe_ok.json +4 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/examples/welcome.json +10 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/schemas/auth_challenge.json +13 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/schemas/auth_response.json +11 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/schemas/broadcast.json +18 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/schemas/bye.json +10 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/schemas/channels.json +10 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/schemas/channels_ok.json +26 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/schemas/custom.json +22 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/schemas/error.json +37 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/schemas/hello.json +59 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/schemas/kick.json +13 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/schemas/kick_ok.json +12 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/schemas/list.json +10 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/schemas/list_ok.json +24 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/schemas/msg.json +29 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/schemas/ping.json +10 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/schemas/pong.json +10 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/schemas/publish.json +23 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/schemas/send.json +22 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/schemas/shutdown.json +10 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/schemas/shutdown_ok.json +10 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/schemas/subscribe.json +15 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/schemas/subscribe_ok.json +11 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/schemas/unsubscribe.json +15 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/schemas/unsubscribe_ok.json +11 -0
- inter_agent_core-0.2.0.data/data/share/inter-agent/spec/schemas/welcome.json +36 -0
- inter_agent_core-0.2.0.dist-info/METADATA +85 -0
- inter_agent_core-0.2.0.dist-info/RECORD +80 -0
- inter_agent_core-0.2.0.dist-info/WHEEL +5 -0
- inter_agent_core-0.2.0.dist-info/entry_points.txt +10 -0
- inter_agent_core-0.2.0.dist-info/licenses/LICENSE.md +21 -0
- inter_agent_core-0.2.0.dist-info/top_level.txt +1 -0
|
@@ -0,0 +1,131 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import argparse
|
|
4
|
+
import asyncio
|
|
5
|
+
import json
|
|
6
|
+
import uuid
|
|
7
|
+
from collections.abc import Sequence
|
|
8
|
+
from dataclasses import dataclass
|
|
9
|
+
from pathlib import Path
|
|
10
|
+
|
|
11
|
+
import websockets
|
|
12
|
+
|
|
13
|
+
from inter_agent.core.auth import AuthError, AuthProtocolError, client_handshake
|
|
14
|
+
from inter_agent.core.shared import control_hello, resolve_endpoint, resolve_shared_secret
|
|
15
|
+
from inter_agent.core.transport import client_ssl_context, websocket_uri
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
@dataclass(frozen=True)
|
|
19
|
+
class ChannelInfo:
|
|
20
|
+
"""Diagnostic information for a single pub/sub channel."""
|
|
21
|
+
|
|
22
|
+
name: str
|
|
23
|
+
subscribers: tuple[str, ...]
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
@dataclass(frozen=True)
|
|
27
|
+
class ChannelsResult:
|
|
28
|
+
"""Structured channels command result plus the raw protocol response."""
|
|
29
|
+
|
|
30
|
+
raw_response: str
|
|
31
|
+
response: dict[str, object]
|
|
32
|
+
channels: tuple[ChannelInfo, ...]
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
def _text_frame(frame: str | bytes) -> str:
|
|
36
|
+
if isinstance(frame, bytes):
|
|
37
|
+
return frame.decode("utf-8")
|
|
38
|
+
return frame
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
def _json_object(raw: str) -> dict[str, object]:
|
|
42
|
+
payload: object = json.loads(raw)
|
|
43
|
+
if not isinstance(payload, dict):
|
|
44
|
+
raise ValueError("server response must be a JSON object")
|
|
45
|
+
return {str(key): value for key, value in payload.items()}
|
|
46
|
+
|
|
47
|
+
|
|
48
|
+
def _parse_channels(response: dict[str, object]) -> tuple[ChannelInfo, ...]:
|
|
49
|
+
channels = response.get("channels")
|
|
50
|
+
if not isinstance(channels, list):
|
|
51
|
+
raise ValueError("channels response must include channels")
|
|
52
|
+
|
|
53
|
+
result: list[ChannelInfo] = []
|
|
54
|
+
for entry in channels:
|
|
55
|
+
if not isinstance(entry, dict):
|
|
56
|
+
raise ValueError("channels entries must be objects")
|
|
57
|
+
name = entry.get("name")
|
|
58
|
+
subscribers = entry.get("subscribers")
|
|
59
|
+
if not isinstance(name, str):
|
|
60
|
+
raise ValueError("channels entries must include string name")
|
|
61
|
+
if not isinstance(subscribers, list):
|
|
62
|
+
raise ValueError("channels entries must include subscribers list")
|
|
63
|
+
parsed_subscribers: list[str] = []
|
|
64
|
+
for subscriber in subscribers:
|
|
65
|
+
if not isinstance(subscriber, str):
|
|
66
|
+
raise ValueError("channel subscribers must be strings")
|
|
67
|
+
parsed_subscribers.append(subscriber)
|
|
68
|
+
result.append(ChannelInfo(name=name, subscribers=tuple(parsed_subscribers)))
|
|
69
|
+
return tuple(result)
|
|
70
|
+
|
|
71
|
+
|
|
72
|
+
async def list_channels(
|
|
73
|
+
host: str,
|
|
74
|
+
port: int,
|
|
75
|
+
*,
|
|
76
|
+
tls: bool = False,
|
|
77
|
+
data_dir: Path | None = None,
|
|
78
|
+
tls_cert_path: Path | None = None,
|
|
79
|
+
) -> ChannelsResult:
|
|
80
|
+
"""Return pub/sub channel diagnostics through a control connection."""
|
|
81
|
+
secret = resolve_shared_secret().secret
|
|
82
|
+
ssl_context = client_ssl_context(tls, data_dir, tls_cert_path)
|
|
83
|
+
async with websockets.connect(websocket_uri(host, port, tls), ssl=ssl_context) as ws:
|
|
84
|
+
try:
|
|
85
|
+
_ = await client_handshake(ws, secret, control_hello(f"ctl-{uuid.uuid4()}"))
|
|
86
|
+
except AuthError as exc:
|
|
87
|
+
raise SystemExit(str(exc)) from exc
|
|
88
|
+
except (AuthProtocolError, json.JSONDecodeError, UnicodeDecodeError) as exc:
|
|
89
|
+
raise SystemExit(f"server protocol mismatch: {exc}") from exc
|
|
90
|
+
await ws.send(json.dumps({"op": "channels"}))
|
|
91
|
+
raw_response = _text_frame(await ws.recv())
|
|
92
|
+
response = _json_object(raw_response)
|
|
93
|
+
channels = _parse_channels(response) if response.get("op") == "channels_ok" else ()
|
|
94
|
+
return ChannelsResult(
|
|
95
|
+
raw_response=raw_response,
|
|
96
|
+
response=response,
|
|
97
|
+
channels=channels,
|
|
98
|
+
)
|
|
99
|
+
|
|
100
|
+
|
|
101
|
+
def build_parser() -> argparse.ArgumentParser:
|
|
102
|
+
parser = argparse.ArgumentParser(prog="inter-agent-channels")
|
|
103
|
+
parser.add_argument("--host")
|
|
104
|
+
parser.add_argument("--port", type=int)
|
|
105
|
+
parser.add_argument("--tls", dest="tls", action="store_true", default=None)
|
|
106
|
+
parser.add_argument("--no-tls", dest="tls", action="store_false")
|
|
107
|
+
parser.add_argument("--tls-cert")
|
|
108
|
+
return parser
|
|
109
|
+
|
|
110
|
+
|
|
111
|
+
def main(argv: Sequence[str] | None = None) -> int:
|
|
112
|
+
parser = build_parser()
|
|
113
|
+
args = parser.parse_args(argv)
|
|
114
|
+
endpoint = resolve_endpoint(
|
|
115
|
+
args.host, args.port, allow_discovery=True, tls=args.tls, tls_cert_path=args.tls_cert
|
|
116
|
+
)
|
|
117
|
+
result = asyncio.run(
|
|
118
|
+
list_channels(
|
|
119
|
+
endpoint.host,
|
|
120
|
+
endpoint.port,
|
|
121
|
+
tls=endpoint.tls,
|
|
122
|
+
data_dir=endpoint.data_dir,
|
|
123
|
+
tls_cert_path=endpoint.tls_cert_path,
|
|
124
|
+
)
|
|
125
|
+
)
|
|
126
|
+
print(result.raw_response)
|
|
127
|
+
return 0 if result.response.get("op") == "channels_ok" else 1
|
|
128
|
+
|
|
129
|
+
|
|
130
|
+
if __name__ == "__main__":
|
|
131
|
+
raise SystemExit(main())
|
|
@@ -0,0 +1,345 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import argparse
|
|
4
|
+
import asyncio
|
|
5
|
+
import json
|
|
6
|
+
import os
|
|
7
|
+
import sys
|
|
8
|
+
import uuid
|
|
9
|
+
from collections.abc import AsyncGenerator, AsyncIterator, Sequence
|
|
10
|
+
from pathlib import Path
|
|
11
|
+
from typing import TextIO
|
|
12
|
+
|
|
13
|
+
import websockets
|
|
14
|
+
from websockets.asyncio.client import ClientConnection
|
|
15
|
+
|
|
16
|
+
from inter_agent.core.auth import AuthError, AuthProtocolError, client_handshake
|
|
17
|
+
from inter_agent.core.auth import build_hello as build_auth_hello
|
|
18
|
+
from inter_agent.core.shared import resolve_endpoint, resolve_shared_secret
|
|
19
|
+
from inter_agent.core.transport import client_ssl_context, websocket_uri
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
def build_hello(
|
|
23
|
+
session_id: str, name: str, label: str | None = None, client_nonce: str | None = None
|
|
24
|
+
) -> dict[str, object]:
|
|
25
|
+
return build_auth_hello(
|
|
26
|
+
role="agent",
|
|
27
|
+
session_id=session_id,
|
|
28
|
+
name=name,
|
|
29
|
+
label=label,
|
|
30
|
+
capabilities={},
|
|
31
|
+
client_nonce=client_nonce,
|
|
32
|
+
)
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
def _text_frame(frame: str | bytes) -> str:
|
|
36
|
+
if isinstance(frame, bytes):
|
|
37
|
+
return frame.decode("utf-8")
|
|
38
|
+
return frame
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
def _json_object(raw: str) -> dict[str, object]:
|
|
42
|
+
payload: object = json.loads(raw)
|
|
43
|
+
if not isinstance(payload, dict):
|
|
44
|
+
raise ValueError("server response must be a JSON object")
|
|
45
|
+
return {str(key): value for key, value in payload.items()}
|
|
46
|
+
|
|
47
|
+
|
|
48
|
+
async def iter_client_frames(
|
|
49
|
+
host: str,
|
|
50
|
+
port: int,
|
|
51
|
+
name: str,
|
|
52
|
+
label: str | None = None,
|
|
53
|
+
*,
|
|
54
|
+
tls: bool = False,
|
|
55
|
+
data_dir: Path | None = None,
|
|
56
|
+
tls_cert_path: Path | None = None,
|
|
57
|
+
) -> AsyncGenerator[str, None]:
|
|
58
|
+
"""Connect an agent session and yield raw server JSON frames.
|
|
59
|
+
|
|
60
|
+
The first yielded frame is the server welcome response. Subsequent frames are
|
|
61
|
+
peer messages or protocol responses received for the connected session.
|
|
62
|
+
"""
|
|
63
|
+
secret = resolve_shared_secret().secret
|
|
64
|
+
session_id = os.getenv("INTER_AGENT_SESSION_ID", str(uuid.uuid4()))
|
|
65
|
+
hello = build_hello(session_id, name, label)
|
|
66
|
+
ssl_context = client_ssl_context(tls, data_dir, tls_cert_path)
|
|
67
|
+
async with websockets.connect(websocket_uri(host, port, tls), ssl=ssl_context) as ws:
|
|
68
|
+
try:
|
|
69
|
+
yield await client_handshake(ws, secret, hello)
|
|
70
|
+
except AuthError as exc:
|
|
71
|
+
raise SystemExit(str(exc)) from exc
|
|
72
|
+
except (AuthProtocolError, json.JSONDecodeError, UnicodeDecodeError) as exc:
|
|
73
|
+
raise SystemExit(f"server protocol mismatch: {exc}") from exc
|
|
74
|
+
async for msg in ws:
|
|
75
|
+
yield _text_frame(msg)
|
|
76
|
+
|
|
77
|
+
|
|
78
|
+
async def run_client(
|
|
79
|
+
host: str,
|
|
80
|
+
port: int,
|
|
81
|
+
name: str,
|
|
82
|
+
label: str | None = None,
|
|
83
|
+
output: TextIO | None = None,
|
|
84
|
+
*,
|
|
85
|
+
tls: bool = False,
|
|
86
|
+
data_dir: Path | None = None,
|
|
87
|
+
tls_cert_path: Path | None = None,
|
|
88
|
+
) -> None:
|
|
89
|
+
"""Run the connect command behavior using typed inputs instead of argv."""
|
|
90
|
+
stream = output or sys.stdout
|
|
91
|
+
async for msg in iter_client_frames(
|
|
92
|
+
host, port, name, label, tls=tls, data_dir=data_dir, tls_cert_path=tls_cert_path
|
|
93
|
+
):
|
|
94
|
+
print(msg, file=stream)
|
|
95
|
+
|
|
96
|
+
|
|
97
|
+
def build_parser() -> argparse.ArgumentParser:
|
|
98
|
+
parser = argparse.ArgumentParser(prog="inter-agent-connect")
|
|
99
|
+
parser.add_argument("name", nargs="?")
|
|
100
|
+
parser.add_argument("--name", dest="name_option")
|
|
101
|
+
parser.add_argument("--label")
|
|
102
|
+
parser.add_argument("--host")
|
|
103
|
+
parser.add_argument("--port", type=int)
|
|
104
|
+
parser.add_argument("--tls", dest="tls", action="store_true", default=None)
|
|
105
|
+
parser.add_argument("--no-tls", dest="tls", action="store_false")
|
|
106
|
+
parser.add_argument("--tls-cert")
|
|
107
|
+
return parser
|
|
108
|
+
|
|
109
|
+
|
|
110
|
+
def main(argv: Sequence[str] | None = None) -> int:
|
|
111
|
+
parser = build_parser()
|
|
112
|
+
args = parser.parse_args(argv)
|
|
113
|
+
name = args.name_option or args.name
|
|
114
|
+
if not name:
|
|
115
|
+
parser.error("name is required")
|
|
116
|
+
endpoint = resolve_endpoint(
|
|
117
|
+
args.host, args.port, allow_discovery=True, tls=args.tls, tls_cert_path=args.tls_cert
|
|
118
|
+
)
|
|
119
|
+
asyncio.run(
|
|
120
|
+
run_client(
|
|
121
|
+
endpoint.host,
|
|
122
|
+
endpoint.port,
|
|
123
|
+
name,
|
|
124
|
+
args.label,
|
|
125
|
+
tls=endpoint.tls,
|
|
126
|
+
data_dir=endpoint.data_dir,
|
|
127
|
+
tls_cert_path=endpoint.tls_cert_path,
|
|
128
|
+
)
|
|
129
|
+
)
|
|
130
|
+
return 0
|
|
131
|
+
|
|
132
|
+
|
|
133
|
+
_COMMAND_TIMEOUT = 0.1
|
|
134
|
+
_EOF = object()
|
|
135
|
+
|
|
136
|
+
|
|
137
|
+
class _PendingCommand:
|
|
138
|
+
__slots__ = ("expected_ops", "future")
|
|
139
|
+
|
|
140
|
+
def __init__(self, expected_ops: set[str], future: asyncio.Future[str]) -> None:
|
|
141
|
+
self.expected_ops = expected_ops
|
|
142
|
+
self.future = future
|
|
143
|
+
|
|
144
|
+
|
|
145
|
+
class AgentSession:
|
|
146
|
+
"""Persistent async agent connection with channel operations.
|
|
147
|
+
|
|
148
|
+
``async with AgentSession(...) as session:`` establishes a single agent
|
|
149
|
+
session and yields inbound frames through ``async for frame in session:``.
|
|
150
|
+
Subscribe, unsubscribe, and publish operations reuse the same connection
|
|
151
|
+
and session identity.
|
|
152
|
+
"""
|
|
153
|
+
|
|
154
|
+
def __init__(
|
|
155
|
+
self,
|
|
156
|
+
host: str,
|
|
157
|
+
port: int,
|
|
158
|
+
name: str,
|
|
159
|
+
label: str | None = None,
|
|
160
|
+
*,
|
|
161
|
+
tls: bool = False,
|
|
162
|
+
data_dir: Path | None = None,
|
|
163
|
+
tls_cert_path: Path | None = None,
|
|
164
|
+
) -> None:
|
|
165
|
+
self.host = host
|
|
166
|
+
self.port = port
|
|
167
|
+
self.name = name
|
|
168
|
+
self.label = label
|
|
169
|
+
self.tls = tls
|
|
170
|
+
self.data_dir = data_dir
|
|
171
|
+
self.tls_cert_path = tls_cert_path
|
|
172
|
+
|
|
173
|
+
self._ws: ClientConnection | None = None
|
|
174
|
+
self._reader_task: asyncio.Task[None] | None = None
|
|
175
|
+
self._inbox: asyncio.Queue[str | object] = asyncio.Queue()
|
|
176
|
+
self._command_lock = asyncio.Lock()
|
|
177
|
+
self._state_lock = asyncio.Lock()
|
|
178
|
+
self._pending: _PendingCommand | None = None
|
|
179
|
+
self._closed = False
|
|
180
|
+
self._eof_delivered = False
|
|
181
|
+
|
|
182
|
+
async def __aenter__(self) -> AgentSession:
|
|
183
|
+
secret = resolve_shared_secret().secret
|
|
184
|
+
session_id = os.getenv("INTER_AGENT_SESSION_ID", str(uuid.uuid4()))
|
|
185
|
+
hello = build_hello(session_id, self.name, self.label)
|
|
186
|
+
ssl_context = client_ssl_context(self.tls, self.data_dir, self.tls_cert_path)
|
|
187
|
+
self._ws = await websockets.connect(
|
|
188
|
+
websocket_uri(self.host, self.port, self.tls), ssl=ssl_context
|
|
189
|
+
)
|
|
190
|
+
try:
|
|
191
|
+
welcome = await client_handshake(self._ws, secret, hello)
|
|
192
|
+
except AuthError as exc:
|
|
193
|
+
await self._close()
|
|
194
|
+
raise SystemExit(str(exc)) from exc
|
|
195
|
+
except (AuthProtocolError, json.JSONDecodeError, UnicodeDecodeError) as exc:
|
|
196
|
+
await self._close()
|
|
197
|
+
raise SystemExit(f"server protocol mismatch: {exc}") from exc
|
|
198
|
+
await self._inbox.put(welcome)
|
|
199
|
+
self._reader_task = asyncio.create_task(self._reader_loop())
|
|
200
|
+
return self
|
|
201
|
+
|
|
202
|
+
async def __aexit__(
|
|
203
|
+
self,
|
|
204
|
+
exc_type: type[BaseException] | None,
|
|
205
|
+
exc: BaseException | None,
|
|
206
|
+
tb: object | None,
|
|
207
|
+
) -> None:
|
|
208
|
+
await self._close()
|
|
209
|
+
|
|
210
|
+
async def _close(self) -> None:
|
|
211
|
+
self._closed = True
|
|
212
|
+
ws = self._ws
|
|
213
|
+
self._ws = None
|
|
214
|
+
if ws is not None:
|
|
215
|
+
await ws.close()
|
|
216
|
+
task = self._reader_task
|
|
217
|
+
self._reader_task = None
|
|
218
|
+
if task is not None:
|
|
219
|
+
task.cancel()
|
|
220
|
+
try:
|
|
221
|
+
await task
|
|
222
|
+
except asyncio.CancelledError:
|
|
223
|
+
pass
|
|
224
|
+
await self._deliver_eof()
|
|
225
|
+
|
|
226
|
+
async def _deliver_eof(self) -> None:
|
|
227
|
+
async with self._state_lock:
|
|
228
|
+
if self._eof_delivered:
|
|
229
|
+
return
|
|
230
|
+
self._eof_delivered = True
|
|
231
|
+
await self._inbox.put(_EOF)
|
|
232
|
+
|
|
233
|
+
def __aiter__(self) -> AsyncIterator[str]:
|
|
234
|
+
return self._iter_frames()
|
|
235
|
+
|
|
236
|
+
async def _iter_frames(self) -> AsyncGenerator[str, None]:
|
|
237
|
+
while True:
|
|
238
|
+
item = await self._inbox.get()
|
|
239
|
+
if item is _EOF:
|
|
240
|
+
break
|
|
241
|
+
yield item # type: ignore[misc]
|
|
242
|
+
|
|
243
|
+
async def _reader_loop(self) -> None:
|
|
244
|
+
ws = self._ws
|
|
245
|
+
if ws is None:
|
|
246
|
+
return
|
|
247
|
+
try:
|
|
248
|
+
async for msg in ws:
|
|
249
|
+
raw = _text_frame(msg)
|
|
250
|
+
op = self._frame_op(raw)
|
|
251
|
+
routed = False
|
|
252
|
+
async with self._state_lock:
|
|
253
|
+
if self._pending is not None and op in self._pending.expected_ops:
|
|
254
|
+
if not self._pending.future.done():
|
|
255
|
+
self._pending.future.set_result(raw)
|
|
256
|
+
self._pending = None
|
|
257
|
+
routed = True
|
|
258
|
+
if not routed:
|
|
259
|
+
await self._inbox.put(raw)
|
|
260
|
+
except websockets.ConnectionClosed:
|
|
261
|
+
pass
|
|
262
|
+
except asyncio.CancelledError:
|
|
263
|
+
raise
|
|
264
|
+
finally:
|
|
265
|
+
self._closed = True
|
|
266
|
+
async with self._state_lock:
|
|
267
|
+
if self._pending is not None and not self._pending.future.done():
|
|
268
|
+
self._pending.future.set_exception(ConnectionAbortedError("session closed"))
|
|
269
|
+
await self._deliver_eof()
|
|
270
|
+
|
|
271
|
+
@staticmethod
|
|
272
|
+
def _frame_op(raw: str) -> str | None:
|
|
273
|
+
try:
|
|
274
|
+
payload = json.loads(raw)
|
|
275
|
+
except (json.JSONDecodeError, UnicodeDecodeError):
|
|
276
|
+
return None
|
|
277
|
+
if not isinstance(payload, dict):
|
|
278
|
+
return None
|
|
279
|
+
op = payload.get("op")
|
|
280
|
+
return op if isinstance(op, str) else None
|
|
281
|
+
|
|
282
|
+
async def _exchange(
|
|
283
|
+
self,
|
|
284
|
+
payload: dict[str, object],
|
|
285
|
+
expected_ops: set[str],
|
|
286
|
+
*,
|
|
287
|
+
allow_timeout: bool = False,
|
|
288
|
+
) -> str | None:
|
|
289
|
+
async with self._command_lock:
|
|
290
|
+
future: asyncio.Future[str] = asyncio.get_running_loop().create_future()
|
|
291
|
+
async with self._state_lock:
|
|
292
|
+
if self._closed:
|
|
293
|
+
raise RuntimeError("session closed")
|
|
294
|
+
self._pending = _PendingCommand(expected_ops, future)
|
|
295
|
+
try:
|
|
296
|
+
ws = self._ws
|
|
297
|
+
if ws is None:
|
|
298
|
+
raise RuntimeError("session closed")
|
|
299
|
+
await ws.send(json.dumps(payload))
|
|
300
|
+
if allow_timeout:
|
|
301
|
+
try:
|
|
302
|
+
return await asyncio.wait_for(future, timeout=_COMMAND_TIMEOUT)
|
|
303
|
+
except TimeoutError:
|
|
304
|
+
return None
|
|
305
|
+
return await asyncio.wait_for(future, timeout=_COMMAND_TIMEOUT)
|
|
306
|
+
finally:
|
|
307
|
+
async with self._state_lock:
|
|
308
|
+
if self._pending is not None and self._pending.future is future:
|
|
309
|
+
self._pending = None
|
|
310
|
+
|
|
311
|
+
async def subscribe(self, channel: str) -> dict[str, object]:
|
|
312
|
+
"""Subscribe to a channel and return the server's response."""
|
|
313
|
+
raw = await self._exchange(
|
|
314
|
+
{"op": "subscribe", "channel": channel}, {"subscribe_ok", "error"}
|
|
315
|
+
)
|
|
316
|
+
assert raw is not None
|
|
317
|
+
return _json_object(raw)
|
|
318
|
+
|
|
319
|
+
async def unsubscribe(self, channel: str) -> dict[str, object]:
|
|
320
|
+
"""Unsubscribe from a channel and return the server's response."""
|
|
321
|
+
raw = await self._exchange(
|
|
322
|
+
{"op": "unsubscribe", "channel": channel}, {"unsubscribe_ok", "error"}
|
|
323
|
+
)
|
|
324
|
+
assert raw is not None
|
|
325
|
+
return _json_object(raw)
|
|
326
|
+
|
|
327
|
+
async def publish(
|
|
328
|
+
self, channel: str, text: str, from_name: str | None = None
|
|
329
|
+
) -> dict[str, object] | None:
|
|
330
|
+
"""Publish a message to a channel.
|
|
331
|
+
|
|
332
|
+
Returns a parsed protocol error if one is received within the command
|
|
333
|
+
timeout, otherwise ``None`` on success.
|
|
334
|
+
"""
|
|
335
|
+
payload: dict[str, object] = {"op": "publish", "channel": channel, "text": text}
|
|
336
|
+
if from_name is not None:
|
|
337
|
+
payload["from_name"] = from_name
|
|
338
|
+
raw = await self._exchange(payload, {"error"}, allow_timeout=True)
|
|
339
|
+
if raw is None:
|
|
340
|
+
return None
|
|
341
|
+
return _json_object(raw)
|
|
342
|
+
|
|
343
|
+
|
|
344
|
+
if __name__ == "__main__":
|
|
345
|
+
raise SystemExit(main())
|