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.
- aiohttp_tiny_mcp/__init__.py +76 -0
- aiohttp_tiny_mcp/adapter.py +321 -0
- aiohttp_tiny_mcp/auth.py +129 -0
- aiohttp_tiny_mcp/client.py +111 -0
- aiohttp_tiny_mcp/client_base.py +295 -0
- aiohttp_tiny_mcp/console/__init__.py +98 -0
- aiohttp_tiny_mcp/console/console.css +525 -0
- aiohttp_tiny_mcp/console/console.js +1273 -0
- aiohttp_tiny_mcp/console/index.html +100 -0
- aiohttp_tiny_mcp/core.py +327 -0
- aiohttp_tiny_mcp/dispatcher.py +309 -0
- aiohttp_tiny_mcp/endpoint.py +531 -0
- aiohttp_tiny_mcp/exchange.py +279 -0
- aiohttp_tiny_mcp/http_sse.py +267 -0
- aiohttp_tiny_mcp/hub.py +109 -0
- aiohttp_tiny_mcp/models.py +346 -0
- aiohttp_tiny_mcp/namespaces.py +36 -0
- aiohttp_tiny_mcp/postgres.py +454 -0
- aiohttp_tiny_mcp/protocol/__init__.py +0 -0
- aiohttp_tiny_mcp/protocol/selection.py +92 -0
- aiohttp_tiny_mcp/protocol/v2024_11_05.py +30 -0
- aiohttp_tiny_mcp/protocol/v2025_03_26.py +165 -0
- aiohttp_tiny_mcp/protocol/v2025_06_18.py +11 -0
- aiohttp_tiny_mcp/protocol/v2025_11_25.py +164 -0
- aiohttp_tiny_mcp/protocol/v2026_07_28.py +363 -0
- aiohttp_tiny_mcp/py.typed +0 -0
- aiohttp_tiny_mcp/redis.py +195 -0
- aiohttp_tiny_mcp/registry.py +162 -0
- aiohttp_tiny_mcp/request_state.py +107 -0
- aiohttp_tiny_mcp/schema.py +131 -0
- aiohttp_tiny_mcp/sessions.py +347 -0
- aiohttp_tiny_mcp/specs.py +268 -0
- aiohttp_tiny_mcp/sqlite.py +236 -0
- aiohttp_tiny_mcp/sse.py +260 -0
- aiohttp_tiny_mcp/stdio.py +187 -0
- aiohttp_tiny_mcp/stdio_client.py +91 -0
- aiohttp_tiny_mcp/subscriptions.py +124 -0
- aiohttp_tiny_mcp/testing.py +200 -0
- aiohttp_tiny_mcp-0.1.0.dist-info/METADATA +12 -0
- aiohttp_tiny_mcp-0.1.0.dist-info/RECORD +41 -0
- 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"]
|
aiohttp_tiny_mcp/hub.py
ADDED
|
@@ -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"]
|