aiohttp-tiny-mcp 0.1.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.
Files changed (41) hide show
  1. aiohttp_tiny_mcp/__init__.py +76 -0
  2. aiohttp_tiny_mcp/adapter.py +321 -0
  3. aiohttp_tiny_mcp/auth.py +129 -0
  4. aiohttp_tiny_mcp/client.py +111 -0
  5. aiohttp_tiny_mcp/client_base.py +295 -0
  6. aiohttp_tiny_mcp/console/__init__.py +98 -0
  7. aiohttp_tiny_mcp/console/console.css +525 -0
  8. aiohttp_tiny_mcp/console/console.js +1273 -0
  9. aiohttp_tiny_mcp/console/index.html +100 -0
  10. aiohttp_tiny_mcp/core.py +327 -0
  11. aiohttp_tiny_mcp/dispatcher.py +309 -0
  12. aiohttp_tiny_mcp/endpoint.py +531 -0
  13. aiohttp_tiny_mcp/exchange.py +279 -0
  14. aiohttp_tiny_mcp/http_sse.py +267 -0
  15. aiohttp_tiny_mcp/hub.py +109 -0
  16. aiohttp_tiny_mcp/models.py +346 -0
  17. aiohttp_tiny_mcp/namespaces.py +36 -0
  18. aiohttp_tiny_mcp/postgres.py +454 -0
  19. aiohttp_tiny_mcp/protocol/__init__.py +0 -0
  20. aiohttp_tiny_mcp/protocol/selection.py +92 -0
  21. aiohttp_tiny_mcp/protocol/v2024_11_05.py +30 -0
  22. aiohttp_tiny_mcp/protocol/v2025_03_26.py +165 -0
  23. aiohttp_tiny_mcp/protocol/v2025_06_18.py +11 -0
  24. aiohttp_tiny_mcp/protocol/v2025_11_25.py +164 -0
  25. aiohttp_tiny_mcp/protocol/v2026_07_28.py +363 -0
  26. aiohttp_tiny_mcp/py.typed +0 -0
  27. aiohttp_tiny_mcp/redis.py +195 -0
  28. aiohttp_tiny_mcp/registry.py +162 -0
  29. aiohttp_tiny_mcp/request_state.py +107 -0
  30. aiohttp_tiny_mcp/schema.py +131 -0
  31. aiohttp_tiny_mcp/sessions.py +347 -0
  32. aiohttp_tiny_mcp/specs.py +268 -0
  33. aiohttp_tiny_mcp/sqlite.py +236 -0
  34. aiohttp_tiny_mcp/sse.py +260 -0
  35. aiohttp_tiny_mcp/stdio.py +187 -0
  36. aiohttp_tiny_mcp/stdio_client.py +91 -0
  37. aiohttp_tiny_mcp/subscriptions.py +124 -0
  38. aiohttp_tiny_mcp/testing.py +200 -0
  39. aiohttp_tiny_mcp-0.1.0.dist-info/METADATA +12 -0
  40. aiohttp_tiny_mcp-0.1.0.dist-info/RECORD +41 -0
  41. aiohttp_tiny_mcp-0.1.0.dist-info/WHEEL +4 -0
@@ -0,0 +1,279 @@
1
+ """Per-request context: dependency resolution, progress, MRTR access."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import asyncio
6
+ import inspect
7
+ import json
8
+ import logging
9
+ from collections.abc import Awaitable, Callable, Mapping
10
+ from contextlib import AsyncExitStack, asynccontextmanager
11
+ from dataclasses import dataclass
12
+ from typing import TYPE_CHECKING, Any
13
+
14
+ from aiohttp import web
15
+
16
+ from .adapter import Adapter, RequestLike
17
+ from .auth import Principal
18
+ from .core import (
19
+ Answer,
20
+ AnswerAction,
21
+ Call,
22
+ ClientInfo,
23
+ InputRequest,
24
+ NeedInput,
25
+ answer_of,
26
+ logs_at,
27
+ )
28
+ from .hub import ASK, Hub, topic
29
+ from .sessions import Session, SessionAccess
30
+ from .sessions import new_session_id as new_id
31
+ from .sse import SSEResponse
32
+
33
+ if TYPE_CHECKING:
34
+ from .registry import Registry
35
+
36
+ log = logging.getLogger("aiohttp_tiny_mcp")
37
+
38
+
39
+ def is_reply(body: Any) -> bool:
40
+ """Pushed-question answers are JSON-RPC responses without a method."""
41
+ return (
42
+ isinstance(body, dict)
43
+ and "method" not in body
44
+ and ("result" in body or "error" in body)
45
+ and isinstance(body.get("id"), str)
46
+ )
47
+
48
+
49
+ async def relay_reply(hub: Hub, body: Mapping[str, Any]) -> None:
50
+ """Publish an answer to its namespaced question topic, which may be read on another node."""
51
+ reply = body["result"] if "result" in body else body["error"]
52
+ await hub.publish(topic(ASK, str(body["id"])), {"reply": reply})
53
+
54
+
55
+ @dataclass(frozen=True, slots=True)
56
+ class Instance:
57
+ """An existing dependency object registered with Registry.provide_instance."""
58
+
59
+ value: Any
60
+
61
+
62
+ class Exchange:
63
+ def __init__(
64
+ self,
65
+ registry: Registry,
66
+ request: RequestLike,
67
+ adapter: Adapter,
68
+ call: Call,
69
+ session: Session | None = None,
70
+ ) -> None:
71
+ self.registry = registry
72
+ self.request = request
73
+ self.adapter = adapter
74
+ self.call = call
75
+ #: Application values kept between calls, or None where this request
76
+ #: reached no session.
77
+ self.session = session
78
+ self.principal: Principal | None = None
79
+ self.keep_log_level: Callable[[str], None] | None = None
80
+ self.sse: SSEResponse | None = None
81
+ self.send: Callable[[Mapping[str, Any]], Awaitable[None]] | None = None
82
+ self.stack = AsyncExitStack()
83
+ self.resolved: dict[type, Any] = {}
84
+ self.cancelled = asyncio.Event()
85
+
86
+ @asynccontextmanager
87
+ async def scope(self):
88
+ """Release per-call dependencies, propagating handler exceptions into providers for
89
+ rollback.
90
+ """
91
+ async with self.stack:
92
+ yield
93
+
94
+ async def resolve(self, kind: type) -> Any:
95
+ if kind is Exchange:
96
+ return self
97
+ if kind is Principal:
98
+ return self.principal
99
+ if kind in self.resolved:
100
+ return self.resolved[kind]
101
+ source = self.registry.providers.get(kind)
102
+ if source is None:
103
+ raise LookupError(f"no provider for {kind!r}")
104
+ if isinstance(source, Instance):
105
+ value = source.value
106
+ elif isinstance(source, web.AppKey):
107
+ value = self.request.app[source]
108
+ elif inspect.isasyncgenfunction(source):
109
+ value = await self.stack.enter_async_context(asynccontextmanager(source)(self))
110
+ else:
111
+ value = await source(self)
112
+ self.resolved[kind] = value
113
+ return value
114
+
115
+ @property
116
+ def id(self) -> str | int:
117
+ assert self.call.id is not None
118
+ return self.call.id
119
+
120
+ @property
121
+ def can_ask(self) -> bool:
122
+ """Whether input can be requested: MRTR requires declared elicitation; the tool-argument
123
+ fallback does not.
124
+ """
125
+ if self.adapter.asks_in_arguments:
126
+ return True
127
+ return self.asks_somehow and self.declared_elicitation
128
+
129
+ @property
130
+ def asks_somehow(self) -> bool:
131
+ """Whether this revision can put a question to a client at all."""
132
+ return self.adapter.can_ask or self.adapter.can_push_ask or self.adapter.asks_in_arguments
133
+
134
+ @property
135
+ def declared_elicitation(self) -> bool:
136
+ """Capability presence is sufficient; an empty mapping is a valid declaration."""
137
+ return self.call.client.capabilities.get("elicitation") is not None
138
+
139
+ @property
140
+ def sessions(self) -> SessionAccess:
141
+ """Access sessions by explicit handles, including on revisions without protocol sessions."""
142
+ return SessionAccess(self.registry.session_store, self.registry.session_ttl_seconds)
143
+
144
+ @property
145
+ def answers(self) -> dict[str, Any]:
146
+ return dict(self.call.answers)
147
+
148
+ def answered(self, key: str) -> bool:
149
+ """Check for a reply, including accepted forms with empty content."""
150
+ return key in self.call.actions or key in self.call.answers
151
+
152
+ def accepted(self, key: str) -> bool:
153
+ """Check the action; accepted and declined replies can both have empty content."""
154
+ return self.call.actions.get(key) is AnswerAction.ACCEPT
155
+
156
+ async def ask(
157
+ self,
158
+ key: str,
159
+ request: InputRequest,
160
+ *,
161
+ default: Mapping[str, Any] | None = None,
162
+ ) -> Answer:
163
+ """Ask for input and return the answer.
164
+
165
+ MRTR raises NeedInput and restarts the handler when the answer arrives, possibly on
166
+ another node. Put irreversible work after the final ask; preceding work may run again.
167
+
168
+ For clients that cannot be asked, `default` accepts an elicit_accept/decline/cancel
169
+ result. Without a default, the call fails.
170
+ """
171
+ if self.answered(key):
172
+ return Answer(
173
+ action=self.call.actions.get(key, AnswerAction.ACCEPT),
174
+ content=self.call.answers.get(key) or {},
175
+ )
176
+ if self.can_ask and self.adapter.can_push_ask:
177
+ return await self.push_ask(key, request)
178
+ if default is not None and not self.can_ask:
179
+ return answer_of(default)
180
+ raise NeedInput({key: request})
181
+
182
+ async def push_ask(self, key: str, request: InputRequest) -> Answer:
183
+ """Send a question on the active stream and wait through the hub.
184
+
185
+ The reply may reach another node. Capture the cursor before sending so fast replies are
186
+ not missed.
187
+ """
188
+ hub = self.registry.hub
189
+ wire_id = new_id()
190
+ where = topic(ASK, wire_id)
191
+ cursor = await hub.position(where)
192
+ await self.emit({"jsonrpc": "2.0", "id": wire_id, **request})
193
+ loop = asyncio.get_running_loop()
194
+ deadline = loop.time() + self.registry.ask_timeout_seconds
195
+ try:
196
+ while not self.cancelled.is_set():
197
+ left = deadline - loop.time()
198
+ if left <= 0:
199
+ return Answer(action=AnswerAction.CANCEL)
200
+ messages, cursor = await hub.poll(
201
+ where, cursor, timeout=min(left, self.registry.hub_poll_seconds)
202
+ )
203
+ for message in messages:
204
+ reply = message.get("reply")
205
+ if isinstance(reply, Mapping):
206
+ return answer_of(reply)
207
+ return Answer(action=AnswerAction.CANCEL)
208
+ finally:
209
+ await hub.delete(where)
210
+
211
+ def action(self, key: str) -> AnswerAction | None:
212
+ """Return ACCEPT, DECLINE, CANCEL, or None if unanswered."""
213
+ return self.call.actions.get(key)
214
+
215
+ @property
216
+ def state(self) -> Any:
217
+ return self.call.state
218
+
219
+ @property
220
+ def client_info(self) -> ClientInfo:
221
+ return self.call.client
222
+
223
+ @property
224
+ def log_level(self) -> str | None:
225
+ return self.call.log_level
226
+
227
+ @property
228
+ def progress_token(self) -> str | int:
229
+ return self.call.progress_token if self.call.progress_token is not None else self.id
230
+
231
+ async def emit(self, payload: Mapping[str, Any]) -> None:
232
+ if self.cancelled.is_set():
233
+ return
234
+ if self.send is not None:
235
+ await self.send(payload)
236
+ return
237
+ assert self.sse is not None
238
+ text = json.dumps(payload, ensure_ascii=False)
239
+ log.debug("-> [%s] %s", self.adapter.version, text)
240
+ await self.sse.send(text)
241
+
242
+ def cancel(self) -> None:
243
+ self.cancelled.set()
244
+
245
+ async def wait_cancelled(self) -> None:
246
+ await self.cancelled.wait()
247
+
248
+ def logs(self, level: str) -> bool:
249
+ """Whether a message of `level` would reach this client."""
250
+ return logs_at(self.log_level, level)
251
+
252
+ async def log(self, level: str, data: Any, *, logger: str | None = None) -> None:
253
+ """Emit a message at the requested severity on the current request stream. Without a
254
+ stream, emit nothing.
255
+ """
256
+ if not self.logs(level):
257
+ return
258
+ params: dict[str, Any] = {"level": level, "data": data}
259
+ if logger is not None:
260
+ params["logger"] = logger
261
+ await self.emit({"jsonrpc": "2.0", "method": "notifications/message", "params": params})
262
+
263
+ async def progress(
264
+ self, progress: float, total: float | None = None, message: str | None = None
265
+ ) -> None:
266
+ """Emit on the current HTTP response or stdio request."""
267
+ if self.sse is None and self.send is None:
268
+ return
269
+ params: dict[str, Any] = {"progressToken": self.progress_token, "progress": progress}
270
+ if total is not None:
271
+ params["total"] = total
272
+ if message is not None and self.adapter.progress_message:
273
+ params["message"] = message
274
+ await self.emit({"jsonrpc": "2.0", "method": "notifications/progress", "params": params})
275
+
276
+ def open(self, *, compress: bool = True, **headers: str) -> SSEResponse:
277
+ """The response this call streams on. Prepared by the caller."""
278
+ self.sse = SSEResponse(compress=compress, headers=headers)
279
+ return self.sse
@@ -0,0 +1,267 @@
1
+ """The legacy 2024-11-05 HTTP+SSE transport."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import asyncio
6
+ import json
7
+ import logging
8
+ from contextlib import suppress
9
+ from dataclasses import replace
10
+
11
+ from aiohttp import web
12
+
13
+ from .adapter import Adapter
14
+ from .core import (
15
+ Call,
16
+ DecodeFailure,
17
+ Failure,
18
+ FailureKind,
19
+ Operation,
20
+ Preamble,
21
+ Rejected,
22
+ Value,
23
+ )
24
+ from .dispatcher import Dispatcher
25
+ from .endpoint import Endpoint
26
+ from .exchange import Exchange
27
+ from .hub import topic
28
+ from .namespaces import scoped
29
+ from .protocol.selection import AdapterSet
30
+ from .registry import Registry
31
+ from .sessions import (
32
+ Session,
33
+ SessionRecord,
34
+ handshake_data,
35
+ new_session_id,
36
+ stored_capabilities,
37
+ stored_log_level,
38
+ stored_version,
39
+ )
40
+ from .sse import SSEResponse
41
+
42
+ log = logging.getLogger(__name__)
43
+
44
+ STREAM = "sse"
45
+
46
+ VERSION = "2024-11-05"
47
+
48
+
49
+ class SseEndpoint:
50
+ """Two endpoints: one to listen on, one to send to."""
51
+
52
+ def __init__(
53
+ self,
54
+ registry: Registry,
55
+ *,
56
+ adapters: AdapterSet | None = None,
57
+ allowed_origins: set[str] | None = None,
58
+ trust_proxy_origin_validation: bool = False,
59
+ compress: bool = True,
60
+ sse_path: str = "/sse",
61
+ message_path: str = "/messages",
62
+ ) -> None:
63
+ self.registry = registry
64
+ self.adapters = adapters or AdapterSet.default()
65
+ self.dispatcher = Dispatcher(registry)
66
+ self.origins = Endpoint(
67
+ registry,
68
+ adapters=self.adapters,
69
+ allowed_origins=allowed_origins,
70
+ trust_proxy_origin_validation=trust_proxy_origin_validation,
71
+ )
72
+ self.compress = compress
73
+ self.sse_path = sse_path
74
+ self.message_path = message_path
75
+
76
+ @property
77
+ def adapter(self) -> Adapter:
78
+ """The one that renders a failure before a revision is known."""
79
+ return self.adapters.by_version[VERSION]
80
+
81
+ def speaking(self, pre: Preamble, record: SessionRecord) -> Adapter:
82
+ """Select a handshake-capable revision."""
83
+ chosen = self.adapters.select(pre, stored_version(record) or VERSION)
84
+ if not chosen.has_handshake:
85
+ raise Rejected(
86
+ Failure(
87
+ FailureKind.UNSUPPORTED_VERSION,
88
+ f"{chosen.version} is not spoken over this transport",
89
+ data={
90
+ "supported": [a.version for a in self.adapters.adapters if a.has_handshake]
91
+ },
92
+ )
93
+ )
94
+ return chosen
95
+
96
+ async def listen(self, request: web.Request) -> web.StreamResponse:
97
+ """Open a stream and send its POST address."""
98
+ try:
99
+ self.origins.check_origin(request)
100
+ except Rejected as e:
101
+ return self.origins.render_failure(self.adapter, e.failure)
102
+
103
+ session_id = await self.open_session()
104
+ where = topic(STREAM, session_id)
105
+ hub = self.registry.hub
106
+ cursor = await hub.position(where)
107
+
108
+ response = SSEResponse(compress=self.compress)
109
+ await response.prepare(request)
110
+ posting = self.mounted(request) + self.message_path
111
+ await self.write(response, "endpoint", f"{posting}?session_id={session_id}")
112
+
113
+ try:
114
+ await self.relay_until_disconnect(request, response, where, cursor)
115
+ finally:
116
+ await hub.delete(where)
117
+ await self.registry.session_store.delete(scoped(session_id))
118
+ return response
119
+
120
+ async def open_session(self) -> str:
121
+ """Create the short-lived session represented by an open stream."""
122
+ session_id = new_session_id()
123
+ await self.registry.session_store.create(
124
+ scoped(session_id),
125
+ handshake_data(VERSION, {}),
126
+ ttl_seconds=self.registry.session_ttl_seconds,
127
+ )
128
+ return session_id
129
+
130
+ async def relay_until_disconnect(
131
+ self, request: web.Request, response: SSEResponse, where: str, cursor: str
132
+ ) -> None:
133
+ """Run a relay until its client disconnects, then always join it."""
134
+ relay = asyncio.create_task(self.relay(response, where, cursor))
135
+ try:
136
+ while not relay.done():
137
+ await asyncio.wait({relay}, timeout=0.05)
138
+ transport = request.transport
139
+ if transport is None or transport.is_closing():
140
+ break
141
+ finally:
142
+ relay.cancel()
143
+ with suppress(asyncio.CancelledError, ConnectionError):
144
+ await relay
145
+
146
+ async def relay(self, response: SSEResponse, where: str, cursor: str) -> None:
147
+ """Write everything published for this connection, until cancelled."""
148
+ hub = self.registry.hub
149
+ while True:
150
+ messages, cursor = await hub.poll(where, cursor, timeout=self.registry.hub_poll_seconds)
151
+ for payload in messages:
152
+ await self.write(response, "message", json.dumps(payload, ensure_ascii=False))
153
+
154
+ async def write(self, response: SSEResponse, event: str, data: str) -> None:
155
+ log.debug("-> [%s] %s %s", VERSION, event, data)
156
+ await response.send(data, event=event)
157
+
158
+ async def receive(self, request: web.Request) -> web.StreamResponse:
159
+ """Take one message and answer 202. The reply goes to the stream."""
160
+ try:
161
+ self.origins.check_origin(request)
162
+ except Rejected as e:
163
+ return self.origins.render_failure(self.adapter, e.failure)
164
+
165
+ session_id = request.query.get("session_id", "")
166
+ record = await self.registry.session_store.get(scoped(session_id)) if session_id else None
167
+ if record is None:
168
+ return web.json_response({"error": "no such session"}, status=404)
169
+
170
+ raw = await request.read()
171
+ log.debug("<- [%s] %s", VERSION, raw.decode("utf-8", "replace"))
172
+ where = topic(STREAM, session_id)
173
+
174
+ pre = Preamble.of(raw, request.headers)
175
+ try:
176
+ adapter = self.speaking(pre, record)
177
+ items = adapter.decode(pre)
178
+ except Rejected as e:
179
+ await self.registry.hub.publish(where, self.adapter.encode_failure(None, e.failure))
180
+ return web.Response(status=202)
181
+
182
+ for item in items:
183
+ await self.serve(adapter, item, session_id, record, where)
184
+ return web.Response(status=202)
185
+
186
+ async def serve(
187
+ self,
188
+ adapter: Adapter,
189
+ item: Call | DecodeFailure,
190
+ session_id: str,
191
+ record: SessionRecord,
192
+ where: str,
193
+ ) -> None:
194
+ hub = self.registry.hub
195
+ if isinstance(item, DecodeFailure):
196
+ if item.must_respond:
197
+ await hub.publish(where, adapter.encode_failure(item.id, item.failure))
198
+ return
199
+
200
+ item.client = replace(
201
+ item.client,
202
+ capabilities={**stored_capabilities(record), **item.client.capabilities},
203
+ )
204
+ if item.log_level is None:
205
+ item.log_level = stored_log_level(record)
206
+
207
+ session = Session(
208
+ self.registry.session_store, session_id, record, self.registry.session_ttl_seconds
209
+ )
210
+ ex = Exchange(self.registry, self.origins, adapter, item, session=session)
211
+ ex.send = lambda payload: hub.publish(where, payload)
212
+
213
+ try:
214
+ outcome = await self.dispatcher.run(ex)
215
+ except Exception as e: # noqa: BLE001 -- a failed call must still answer
216
+ log.exception("sse request failed")
217
+ outcome = Failure(FailureKind.INTERNAL, f"{type(e).__name__}: {e}")
218
+
219
+ if item.operation is Operation.DESCRIBE and isinstance(outcome, Value):
220
+ negotiated = getattr(outcome.result, "protocol_version", None)
221
+ if isinstance(negotiated, str):
222
+ await session.remember_version(negotiated)
223
+
224
+ if item.is_notification:
225
+ return
226
+ await hub.publish(where, adapter.encode(item, self.registry, outcome))
227
+
228
+ def mounted(self, request: web.Request) -> str:
229
+ """The prefix the client reached this stream through.
230
+
231
+ A subapplication adds its prefix to the request path and not to the
232
+ route, and the address on the stream is one a client posts to, so it
233
+ has to carry that prefix.
234
+ """
235
+ path = request.path
236
+ return path[: -len(self.sse_path)] if path.endswith(self.sse_path) else ""
237
+
238
+ def routes(
239
+ self, sse_path: str | None = None, message_path: str | None = None
240
+ ) -> list[web.RouteDef]:
241
+ """One route to listen on, one to post to.
242
+
243
+ Either path may be set here or on the constructor. Both are kept,
244
+ because the stream names the posting path to the client.
245
+ """
246
+ if sse_path is not None:
247
+ self.sse_path = sse_path
248
+ if message_path is not None:
249
+ self.message_path = message_path
250
+ log.debug("HTTP+SSE stream at %s, messages at %s", self.sse_path, self.message_path)
251
+ return [
252
+ web.get(self.sse_path, self.listen),
253
+ web.post(self.message_path, self.receive),
254
+ ]
255
+
256
+ def setup(
257
+ self,
258
+ app: web.Application,
259
+ sse_path: str | None = None,
260
+ message_path: str | None = None,
261
+ ) -> web.Application:
262
+ log.debug("adding the HTTP+SSE routes to %r", app)
263
+ app.add_routes(self.routes(sse_path, message_path))
264
+ return app
265
+
266
+
267
+ __all__ = ["STREAM", "VERSION", "SseEndpoint"]
@@ -0,0 +1,109 @@
1
+ """Cursor-based event storage shared by reply and notification streams.
2
+
3
+ Applications supply the backend; MemoryHub serves one process. Shared storage lets one node
4
+ publish replies that another is awaiting.
5
+ """
6
+
7
+ from __future__ import annotations
8
+
9
+ import asyncio
10
+ from abc import abstractmethod
11
+ from collections.abc import Mapping, Sequence
12
+ from typing import Any, Protocol, runtime_checkable
13
+
14
+ from .namespaces import scoped
15
+
16
+ Cursor = str
17
+
18
+ START: Cursor = ""
19
+
20
+ NOTIFICATIONS = "notifications"
21
+
22
+ ASK = "ask"
23
+
24
+
25
+ @runtime_checkable
26
+ class Hub(Protocol):
27
+ """Application-provided event storage.
28
+
29
+ Subclass it to have the methods checked and the missing ones refused, or
30
+ supply any object with these four methods: this is a protocol, so a
31
+ backend that inherits nothing is still a Hub.
32
+ """
33
+
34
+ @abstractmethod
35
+ async def publish(self, topic: str, message: Mapping[str, Any]) -> None:
36
+ """Append one message to `topic`."""
37
+
38
+ @abstractmethod
39
+ async def position(self, topic: str) -> Cursor:
40
+ """Capture the cursor before triggering a publish so replies preceding the first poll are
41
+ included.
42
+ """
43
+
44
+ @abstractmethod
45
+ async def poll(
46
+ self, topic: str, cursor: Cursor, *, timeout: float
47
+ ) -> tuple[Sequence[Mapping[str, Any]], Cursor]:
48
+ """Return messages after cursor and the next cursor.
49
+
50
+ Wait up to timeout seconds for a message. Backends choose whether to poll or wait for
51
+ notification; timeout is a deadline, not a polling interval.
52
+ """
53
+
54
+ @abstractmethod
55
+ async def delete(self, topic: str) -> None:
56
+ """Delete a completed topic."""
57
+
58
+
59
+ def topic(kind: str, name: str | None = None) -> str:
60
+ """Prefix the topic with the current namespace to isolate callers."""
61
+ return scoped(kind if name is None else f"{kind}/{name}")
62
+
63
+
64
+ class MemoryHub(Hub):
65
+ """Single-process hub, waiting on a condition rather than sleeping."""
66
+
67
+ def __init__(self) -> None:
68
+ self.rows: dict[str, list[tuple[int, Mapping[str, Any]]]] = {}
69
+ self.last_id = 0
70
+ self.arrived = asyncio.Condition()
71
+
72
+ async def publish(self, topic: str, message: Mapping[str, Any]) -> None:
73
+ async with self.arrived:
74
+ self.last_id += 1
75
+ self.rows.setdefault(topic, []).append((self.last_id, dict(message)))
76
+ self.arrived.notify_all()
77
+
78
+ async def position(self, topic: str) -> Cursor:
79
+ rows = self.rows.get(topic)
80
+ return str(rows[-1][0]) if rows else START
81
+
82
+ def after(self, topic: str, cursor: Cursor) -> list[tuple[int, Mapping[str, Any]]]:
83
+ least = int(cursor) if cursor else 0
84
+ return [row for row in self.rows.get(topic, []) if row[0] > least]
85
+
86
+ async def poll(
87
+ self, topic: str, cursor: Cursor, *, timeout: float
88
+ ) -> tuple[Sequence[Mapping[str, Any]], Cursor]:
89
+ loop = asyncio.get_running_loop()
90
+ deadline = loop.time() + timeout
91
+ async with self.arrived:
92
+ while True:
93
+ found = self.after(topic, cursor)
94
+ if found:
95
+ return [message for _, message in found], str(found[-1][0])
96
+ left = deadline - loop.time()
97
+ if left <= 0:
98
+ return [], cursor
99
+ try:
100
+ await asyncio.wait_for(self.arrived.wait(), left)
101
+ except asyncio.TimeoutError:
102
+ return [], cursor
103
+
104
+ async def delete(self, topic: str) -> None:
105
+ async with self.arrived:
106
+ self.rows.pop(topic, None)
107
+
108
+
109
+ __all__ = ["ASK", "NOTIFICATIONS", "START", "Cursor", "Hub", "MemoryHub", "topic"]