kimi-agent-module-api 1.0.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.
- kimi_agent_module_api/__init__.py +175 -0
- kimi_agent_module_api/contracts.py +1127 -0
- kimi_agent_module_api/events.py +140 -0
- kimi_agent_module_api/images.py +34 -0
- kimi_agent_module_api/py.typed +1 -0
- kimi_agent_module_api/settings.py +33 -0
- kimi_agent_module_api/testing.py +1036 -0
- kimi_agent_module_api/tools.py +63 -0
- kimi_agent_module_api/trust.py +36 -0
- kimi_agent_module_api-1.0.0.dist-info/METADATA +51 -0
- kimi_agent_module_api-1.0.0.dist-info/RECORD +14 -0
- kimi_agent_module_api-1.0.0.dist-info/WHEEL +5 -0
- kimi_agent_module_api-1.0.0.dist-info/licenses/LICENSE +21 -0
- kimi_agent_module_api-1.0.0.dist-info/top_level.txt +1 -0
|
@@ -0,0 +1,1036 @@
|
|
|
1
|
+
"""Protocol-level fakes for module unit tests.
|
|
2
|
+
|
|
3
|
+
Every fake here satisfies one runtime port from ``contracts`` with plain
|
|
4
|
+
Python and records what a module asked of it. Nothing imports Discord, the
|
|
5
|
+
database, or core runtime packages, so a module package can unit-test its own
|
|
6
|
+
logic with only ``kimi_agent_module_api`` installed. The integration harness
|
|
7
|
+
that composes real core services lives in core's ``modules.testing``.
|
|
8
|
+
"""
|
|
9
|
+
|
|
10
|
+
from __future__ import annotations
|
|
11
|
+
|
|
12
|
+
import asyncio
|
|
13
|
+
import contextlib
|
|
14
|
+
import fnmatch
|
|
15
|
+
import hashlib
|
|
16
|
+
import math
|
|
17
|
+
import sys
|
|
18
|
+
import uuid
|
|
19
|
+
from collections.abc import AsyncIterator, Callable, Mapping, Sequence
|
|
20
|
+
from dataclasses import dataclass, field
|
|
21
|
+
from typing import Any, TypeVar, overload
|
|
22
|
+
|
|
23
|
+
from pydantic_settings import BaseSettings
|
|
24
|
+
|
|
25
|
+
from kimi_agent_module_api.trust import TrustTier
|
|
26
|
+
|
|
27
|
+
from kimi_agent_module_api.contracts import (
|
|
28
|
+
ALL_DISCORD_ACTIONS,
|
|
29
|
+
Backoff,
|
|
30
|
+
ChannelSnapshot,
|
|
31
|
+
CommandSpec,
|
|
32
|
+
ConfigSnapshot,
|
|
33
|
+
Event,
|
|
34
|
+
EventHandler,
|
|
35
|
+
GuildSettingsSnapshot,
|
|
36
|
+
HealthState,
|
|
37
|
+
HostNotAllowed,
|
|
38
|
+
HttpResponse,
|
|
39
|
+
InviteSnapshot,
|
|
40
|
+
JobHandler,
|
|
41
|
+
JobInfo,
|
|
42
|
+
JobRun,
|
|
43
|
+
MemberSnapshot,
|
|
44
|
+
MessagePage,
|
|
45
|
+
MessageRef,
|
|
46
|
+
MessageSnapshot,
|
|
47
|
+
TABLE_NAME_RE,
|
|
48
|
+
MigrationContext,
|
|
49
|
+
ModuleContractError,
|
|
50
|
+
ModuleHealth,
|
|
51
|
+
OutgoingEmbed,
|
|
52
|
+
ProposalActor,
|
|
53
|
+
ProposalError,
|
|
54
|
+
ProposalRef,
|
|
55
|
+
ProposalState,
|
|
56
|
+
RoleSnapshot,
|
|
57
|
+
ScopedModuleMigration,
|
|
58
|
+
ServiceUnavailable,
|
|
59
|
+
TrustTierName,
|
|
60
|
+
UndeclaredDiscordAction,
|
|
61
|
+
build_custom_id,
|
|
62
|
+
validate_publish_topic,
|
|
63
|
+
)
|
|
64
|
+
from kimi_agent_module_api.tools import ModuleToolHandler
|
|
65
|
+
from kimi_agent_module_api import (
|
|
66
|
+
BASELINE_CAPABILITIES,
|
|
67
|
+
ModuleCapabilities,
|
|
68
|
+
ModuleLoadContext,
|
|
69
|
+
)
|
|
70
|
+
|
|
71
|
+
_MAX_FLOAT_LOG = math.log(sys.float_info.max)
|
|
72
|
+
_T = TypeVar("_T")
|
|
73
|
+
|
|
74
|
+
|
|
75
|
+
@dataclass(slots=True)
|
|
76
|
+
class _Closable:
|
|
77
|
+
_on_close: Callable[[], object]
|
|
78
|
+
closed: bool = False
|
|
79
|
+
|
|
80
|
+
def close(self) -> None:
|
|
81
|
+
if not self.closed:
|
|
82
|
+
self.closed = True
|
|
83
|
+
self._on_close()
|
|
84
|
+
|
|
85
|
+
|
|
86
|
+
@dataclass(frozen=True, slots=True)
|
|
87
|
+
class ProposedChange:
|
|
88
|
+
proposal_id: str
|
|
89
|
+
module_name: str
|
|
90
|
+
target: str
|
|
91
|
+
content: str
|
|
92
|
+
summary: str
|
|
93
|
+
actor: ProposalActor
|
|
94
|
+
expected_revision: str | None
|
|
95
|
+
|
|
96
|
+
|
|
97
|
+
class FakeProposals:
|
|
98
|
+
"""Actor-scoped, module-bound fragment proposal fake."""
|
|
99
|
+
|
|
100
|
+
def __init__(
|
|
101
|
+
self,
|
|
102
|
+
module_name: str,
|
|
103
|
+
documents: Mapping[str, str] | None = None,
|
|
104
|
+
*,
|
|
105
|
+
target_guilds: Mapping[str, str] | None = None,
|
|
106
|
+
) -> None:
|
|
107
|
+
self.module_name = module_name
|
|
108
|
+
self.documents = dict(documents or {})
|
|
109
|
+
self.target_guilds = dict(target_guilds or {})
|
|
110
|
+
self.changes: list[ProposedChange] = []
|
|
111
|
+
self.refs: dict[str, ProposalRef] = {}
|
|
112
|
+
self._proposal_guilds: dict[str, str] = {}
|
|
113
|
+
|
|
114
|
+
@staticmethod
|
|
115
|
+
def _revision(content: str) -> str:
|
|
116
|
+
return hashlib.sha256(content.encode("utf-8")).hexdigest()
|
|
117
|
+
|
|
118
|
+
def _require_guild(self, target: str, actor: ProposalActor) -> str:
|
|
119
|
+
actor_guild = str(actor.guild_id or "")
|
|
120
|
+
if not actor_guild.isdecimal() or int(actor_guild) <= 0:
|
|
121
|
+
raise ProposalError("proposal actor must belong to a guild")
|
|
122
|
+
target_guild = self.target_guilds.get(target)
|
|
123
|
+
if target_guild is None and target.startswith("guild:"):
|
|
124
|
+
target_guild = target.split(":", 2)[1]
|
|
125
|
+
if target_guild is None:
|
|
126
|
+
raise ProposalError(f"fake has no guild mapping for {target!r}")
|
|
127
|
+
if target_guild != actor_guild:
|
|
128
|
+
raise ProposalError("proposal target must belong to the actor's guild")
|
|
129
|
+
return actor_guild
|
|
130
|
+
|
|
131
|
+
async def snapshot(self, target: str, *, actor: ProposalActor) -> ConfigSnapshot:
|
|
132
|
+
self._require_guild(target, actor)
|
|
133
|
+
content = self.documents.get(target, "")
|
|
134
|
+
return ConfigSnapshot(target, self._revision(content), content)
|
|
135
|
+
|
|
136
|
+
async def propose(
|
|
137
|
+
self,
|
|
138
|
+
*,
|
|
139
|
+
target: str,
|
|
140
|
+
content: str,
|
|
141
|
+
summary: str,
|
|
142
|
+
actor: ProposalActor,
|
|
143
|
+
expected_revision: str | None = None,
|
|
144
|
+
) -> ProposalRef:
|
|
145
|
+
guild_id = self._require_guild(target, actor)
|
|
146
|
+
current = self.documents.get(target, "")
|
|
147
|
+
if expected_revision is not None and expected_revision != self._revision(current):
|
|
148
|
+
raise ProposalError("configuration changed since it was inspected")
|
|
149
|
+
proposal_id = uuid.uuid4().hex
|
|
150
|
+
change = ProposedChange(
|
|
151
|
+
proposal_id,
|
|
152
|
+
self.module_name,
|
|
153
|
+
target,
|
|
154
|
+
content,
|
|
155
|
+
summary,
|
|
156
|
+
actor,
|
|
157
|
+
expected_revision,
|
|
158
|
+
)
|
|
159
|
+
ref = ProposalRef(proposal_id, target, "pending")
|
|
160
|
+
self.changes.append(change)
|
|
161
|
+
self.refs[proposal_id] = ref
|
|
162
|
+
self._proposal_guilds[proposal_id] = guild_id
|
|
163
|
+
return ref
|
|
164
|
+
|
|
165
|
+
async def get(self, proposal_id: str, *, actor: ProposalActor) -> ProposalRef | None:
|
|
166
|
+
actor_guild = str(actor.guild_id or "")
|
|
167
|
+
if self._proposal_guilds.get(proposal_id) != actor_guild:
|
|
168
|
+
return None
|
|
169
|
+
return self.refs.get(proposal_id)
|
|
170
|
+
|
|
171
|
+
def decide(
|
|
172
|
+
self,
|
|
173
|
+
proposal_id: str,
|
|
174
|
+
state: ProposalState,
|
|
175
|
+
decided_by: str,
|
|
176
|
+
decision_reason: str = "",
|
|
177
|
+
) -> ProposalRef:
|
|
178
|
+
current = self.refs[proposal_id]
|
|
179
|
+
updated = ProposalRef(
|
|
180
|
+
current.proposal_id,
|
|
181
|
+
current.target,
|
|
182
|
+
state,
|
|
183
|
+
current.message,
|
|
184
|
+
decided_by,
|
|
185
|
+
decision_reason,
|
|
186
|
+
)
|
|
187
|
+
self.refs[proposal_id] = updated
|
|
188
|
+
return updated
|
|
189
|
+
|
|
190
|
+
|
|
191
|
+
class FakeEvents:
|
|
192
|
+
"""Records publishes; delivers to subscribers only when ``deliver`` is awaited."""
|
|
193
|
+
|
|
194
|
+
def __init__(self, module_name: str) -> None:
|
|
195
|
+
self.module_name = module_name
|
|
196
|
+
self.published: list[Event] = []
|
|
197
|
+
self._subscriptions: list[tuple[str, EventHandler]] = []
|
|
198
|
+
self._clock = 0.0
|
|
199
|
+
|
|
200
|
+
def publish(self, topic: str, payload: Any) -> None:
|
|
201
|
+
validate_publish_topic(self.module_name, topic)
|
|
202
|
+
self._clock += 1.0
|
|
203
|
+
self.published.append(Event(topic, payload, self.module_name, self._clock))
|
|
204
|
+
|
|
205
|
+
def subscribe(self, pattern: str, handler: EventHandler) -> _Closable:
|
|
206
|
+
entry = (pattern, handler)
|
|
207
|
+
self._subscriptions.append(entry)
|
|
208
|
+
return _Closable(lambda: self._subscriptions.remove(entry))
|
|
209
|
+
|
|
210
|
+
async def deliver(self, topic: str, payload: Any, *, source_module: str = "core") -> int:
|
|
211
|
+
"""Push one event through matching subscribers; returns handler count."""
|
|
212
|
+
self._clock += 1.0
|
|
213
|
+
event = Event(topic, payload, source_module, self._clock)
|
|
214
|
+
count = 0
|
|
215
|
+
for pattern, handler in list(self._subscriptions):
|
|
216
|
+
if fnmatch.fnmatchcase(topic, pattern):
|
|
217
|
+
await handler(event)
|
|
218
|
+
count += 1
|
|
219
|
+
return count
|
|
220
|
+
|
|
221
|
+
@property
|
|
222
|
+
def subscriptions(self) -> tuple[str, ...]:
|
|
223
|
+
return tuple(pattern for pattern, _ in self._subscriptions)
|
|
224
|
+
|
|
225
|
+
|
|
226
|
+
@dataclass(slots=True)
|
|
227
|
+
class _FakeJob:
|
|
228
|
+
key: str
|
|
229
|
+
handler: str
|
|
230
|
+
run_at: float
|
|
231
|
+
interval: float | None
|
|
232
|
+
payload: Mapping[str, Any]
|
|
233
|
+
backoff: Backoff = field(default_factory=Backoff)
|
|
234
|
+
attempt: int = 0
|
|
235
|
+
last_error: str | None = None
|
|
236
|
+
|
|
237
|
+
|
|
238
|
+
def _retry_delay(backoff: Backoff, attempt: int) -> float:
|
|
239
|
+
exponent = max(0, attempt - 1)
|
|
240
|
+
base = backoff.base_seconds
|
|
241
|
+
cap = backoff.max_seconds
|
|
242
|
+
if base >= cap:
|
|
243
|
+
return cap
|
|
244
|
+
if exponent == 0 or backoff.multiplier == 1:
|
|
245
|
+
return base
|
|
246
|
+
|
|
247
|
+
growth_log = exponent * math.log(backoff.multiplier)
|
|
248
|
+
if growth_log >= math.log(cap) - math.log(base):
|
|
249
|
+
return cap
|
|
250
|
+
|
|
251
|
+
if growth_log <= _MAX_FLOAT_LOG:
|
|
252
|
+
try:
|
|
253
|
+
return base * (backoff.multiplier**exponent)
|
|
254
|
+
except OverflowError:
|
|
255
|
+
pass
|
|
256
|
+
return min(cap, math.exp(math.log(base) + growth_log))
|
|
257
|
+
|
|
258
|
+
|
|
259
|
+
class FakeScheduler:
|
|
260
|
+
"""Jobs run only when the test advances time with ``run_due``.
|
|
261
|
+
|
|
262
|
+
Settlement follows the host: a successful one-shot job is deleted, a
|
|
263
|
+
successful periodic job reschedules from completion, and a failed job of
|
|
264
|
+
either kind is kept and retried after its ``Backoff`` delay with the
|
|
265
|
+
attempt count preserved. Jitter is ignored so reschedules are exact.
|
|
266
|
+
"""
|
|
267
|
+
|
|
268
|
+
def __init__(self) -> None:
|
|
269
|
+
self.handlers: dict[str, JobHandler] = {}
|
|
270
|
+
self.jobs: dict[str, _FakeJob] = {}
|
|
271
|
+
self.runs: list[JobRun] = []
|
|
272
|
+
|
|
273
|
+
def register(self, handler_name: str, handler: JobHandler) -> None:
|
|
274
|
+
self.handlers[handler_name] = handler
|
|
275
|
+
|
|
276
|
+
async def run_at(
|
|
277
|
+
self, key: str, when: float, handler_name: str, payload: Mapping[str, Any] | None = None
|
|
278
|
+
) -> None:
|
|
279
|
+
self.jobs[key] = _FakeJob(key, handler_name, when, None, dict(payload or {}))
|
|
280
|
+
|
|
281
|
+
async def run_every(
|
|
282
|
+
self,
|
|
283
|
+
key: str,
|
|
284
|
+
interval_seconds: float,
|
|
285
|
+
handler_name: str,
|
|
286
|
+
payload: Mapping[str, Any] | None = None,
|
|
287
|
+
*,
|
|
288
|
+
jitter_seconds: float = 0.0,
|
|
289
|
+
backoff: Backoff | None = None,
|
|
290
|
+
) -> None:
|
|
291
|
+
self.jobs[key] = _FakeJob(
|
|
292
|
+
key,
|
|
293
|
+
handler_name,
|
|
294
|
+
0.0,
|
|
295
|
+
interval_seconds,
|
|
296
|
+
dict(payload or {}),
|
|
297
|
+
backoff=backoff or Backoff(),
|
|
298
|
+
)
|
|
299
|
+
|
|
300
|
+
async def cancel(self, key: str) -> bool:
|
|
301
|
+
return self.jobs.pop(key, None) is not None
|
|
302
|
+
|
|
303
|
+
async def list(self) -> Sequence[JobInfo]:
|
|
304
|
+
return [
|
|
305
|
+
JobInfo(job.key, job.handler, job.run_at, job.interval, job.attempt, job.last_error)
|
|
306
|
+
for job in self.jobs.values()
|
|
307
|
+
]
|
|
308
|
+
|
|
309
|
+
async def run_due(self, now: float) -> int:
|
|
310
|
+
"""Run every job scheduled at or before ``now``; returns how many ran."""
|
|
311
|
+
ran = 0
|
|
312
|
+
for job in list(self.jobs.values()):
|
|
313
|
+
if job.run_at > now or job.key not in self.jobs:
|
|
314
|
+
continue
|
|
315
|
+
handler = self.handlers.get(job.handler)
|
|
316
|
+
if handler is None:
|
|
317
|
+
job.last_error = f"no handler {job.handler!r}"
|
|
318
|
+
continue
|
|
319
|
+
job.attempt += 1
|
|
320
|
+
run = JobRun(f"fake-{job.key}", job.key, job.payload, job.attempt, job.run_at)
|
|
321
|
+
self.runs.append(run)
|
|
322
|
+
try:
|
|
323
|
+
await handler(run)
|
|
324
|
+
except Exception as exc:
|
|
325
|
+
job.last_error = repr(exc)
|
|
326
|
+
job.run_at = now + _retry_delay(job.backoff, job.attempt)
|
|
327
|
+
else:
|
|
328
|
+
job.last_error = None
|
|
329
|
+
job.attempt = 0
|
|
330
|
+
if job.interval is None:
|
|
331
|
+
self.jobs.pop(job.key, None)
|
|
332
|
+
else:
|
|
333
|
+
job.run_at = now + job.interval
|
|
334
|
+
ran += 1
|
|
335
|
+
return ran
|
|
336
|
+
|
|
337
|
+
|
|
338
|
+
_HEALTH_ORDER: dict[str, int] = {"healthy": 0, "starting": 1, "degraded": 2, "failed": 3}
|
|
339
|
+
|
|
340
|
+
|
|
341
|
+
class FakeHealth:
|
|
342
|
+
"""Records every report.
|
|
343
|
+
|
|
344
|
+
``history`` holds every call as ``(key, health)``. ``current`` is the latest
|
|
345
|
+
unkeyed report, ``keyed`` the latest live report per key, and ``state`` the
|
|
346
|
+
worst across both, which is what the host would show for the module.
|
|
347
|
+
"""
|
|
348
|
+
|
|
349
|
+
def __init__(self) -> None:
|
|
350
|
+
self.history: list[tuple[str | None, ModuleHealth]] = []
|
|
351
|
+
self.reports: list[ModuleHealth] = []
|
|
352
|
+
self.keyed: dict[str, ModuleHealth] = {}
|
|
353
|
+
|
|
354
|
+
def report(
|
|
355
|
+
self,
|
|
356
|
+
state: HealthState,
|
|
357
|
+
detail: str = "",
|
|
358
|
+
metrics: Mapping[str, float] | None = None,
|
|
359
|
+
*,
|
|
360
|
+
key: str | None = None,
|
|
361
|
+
) -> None:
|
|
362
|
+
health = ModuleHealth(state, detail, dict(metrics or {}), float(len(self.history)))
|
|
363
|
+
self.history.append((key, health))
|
|
364
|
+
if key is not None:
|
|
365
|
+
if state == "healthy" and not detail and not metrics:
|
|
366
|
+
self.keyed.pop(key, None)
|
|
367
|
+
else:
|
|
368
|
+
self.keyed[key] = health
|
|
369
|
+
return
|
|
370
|
+
self.reports.append(health)
|
|
371
|
+
|
|
372
|
+
@property
|
|
373
|
+
def current(self) -> ModuleHealth | None:
|
|
374
|
+
return self.reports[-1] if self.reports else None
|
|
375
|
+
|
|
376
|
+
@property
|
|
377
|
+
def state(self) -> HealthState:
|
|
378
|
+
candidates = [*self.keyed.values(), *([self.current] if self.current else [])]
|
|
379
|
+
if not candidates:
|
|
380
|
+
return "healthy"
|
|
381
|
+
return max((health.state for health in candidates), key=_HEALTH_ORDER.__getitem__)
|
|
382
|
+
|
|
383
|
+
|
|
384
|
+
@dataclass(frozen=True, slots=True)
|
|
385
|
+
class DiscordCall:
|
|
386
|
+
action: str
|
|
387
|
+
args: tuple[Any, ...]
|
|
388
|
+
kwargs: Mapping[str, Any]
|
|
389
|
+
|
|
390
|
+
|
|
391
|
+
class FakeDiscordActions:
|
|
392
|
+
"""Enforces the module's declared actions and records every call."""
|
|
393
|
+
|
|
394
|
+
def __init__(self, module_name: str, declared: frozenset[str] | None = None) -> None:
|
|
395
|
+
self.module_name = module_name
|
|
396
|
+
self.declared = ALL_DISCORD_ACTIONS if declared is None else declared
|
|
397
|
+
self.calls: list[DiscordCall] = []
|
|
398
|
+
self.messages: dict[MessageRef, MessageSnapshot] = {}
|
|
399
|
+
self.members: dict[tuple[int, int], MemberSnapshot] = {}
|
|
400
|
+
self.channels: dict[tuple[int, int], ChannelSnapshot] = {}
|
|
401
|
+
self.histories: dict[tuple[int, int], list[MessageSnapshot]] = {}
|
|
402
|
+
self.pins: dict[tuple[int, int], tuple[MessageSnapshot, ...]] = {}
|
|
403
|
+
self.public_threads: dict[tuple[int, int], tuple[ChannelSnapshot, ...]] = {}
|
|
404
|
+
self.roles: dict[int, tuple[RoleSnapshot, ...]] = {}
|
|
405
|
+
self.invites: dict[int, tuple[InviteSnapshot, ...]] = {}
|
|
406
|
+
self.channel_access: dict[tuple[int, int, int], bool] = {}
|
|
407
|
+
self._next_message_id = 1000
|
|
408
|
+
|
|
409
|
+
def _record(self, action: str, *args: Any, **kwargs: Any) -> None:
|
|
410
|
+
if action not in self.declared:
|
|
411
|
+
raise UndeclaredDiscordAction(self.module_name, action)
|
|
412
|
+
self.calls.append(DiscordCall(action, args, kwargs))
|
|
413
|
+
|
|
414
|
+
def calls_for(self, action: str) -> list[DiscordCall]:
|
|
415
|
+
return [call for call in self.calls if call.action == action]
|
|
416
|
+
|
|
417
|
+
async def send_message(
|
|
418
|
+
self,
|
|
419
|
+
channel_id: int,
|
|
420
|
+
content: str | None = None,
|
|
421
|
+
*,
|
|
422
|
+
embed: OutgoingEmbed | None = None,
|
|
423
|
+
reply_to: MessageRef | None = None,
|
|
424
|
+
components: Sequence[Any] = (),
|
|
425
|
+
) -> MessageRef:
|
|
426
|
+
self._record(
|
|
427
|
+
"send_message",
|
|
428
|
+
channel_id,
|
|
429
|
+
content,
|
|
430
|
+
embed=embed,
|
|
431
|
+
reply_to=reply_to,
|
|
432
|
+
components=components,
|
|
433
|
+
)
|
|
434
|
+
self._next_message_id += 1
|
|
435
|
+
return MessageRef(0, channel_id, self._next_message_id)
|
|
436
|
+
|
|
437
|
+
async def send_dm(
|
|
438
|
+
self, user_id: int, content: str, *, embed: OutgoingEmbed | None = None
|
|
439
|
+
) -> bool:
|
|
440
|
+
self._record("send_dm", user_id, content, embed=embed)
|
|
441
|
+
return True
|
|
442
|
+
|
|
443
|
+
async def edit_message(
|
|
444
|
+
self, ref: MessageRef, content: str | None = None, *, embed: OutgoingEmbed | None = None
|
|
445
|
+
) -> None:
|
|
446
|
+
self._record("edit_message", ref, content, embed=embed)
|
|
447
|
+
|
|
448
|
+
async def delete_message(self, ref: MessageRef, *, reason: str = "") -> None:
|
|
449
|
+
self._record("delete_message", ref, reason=reason)
|
|
450
|
+
self.messages.pop(ref, None)
|
|
451
|
+
|
|
452
|
+
async def ban(
|
|
453
|
+
self,
|
|
454
|
+
guild_id: int,
|
|
455
|
+
user_id: int,
|
|
456
|
+
*,
|
|
457
|
+
actor_id: int | None,
|
|
458
|
+
reason: str,
|
|
459
|
+
delete_message_seconds: int = 0,
|
|
460
|
+
) -> None:
|
|
461
|
+
self._record(
|
|
462
|
+
"ban",
|
|
463
|
+
guild_id,
|
|
464
|
+
user_id,
|
|
465
|
+
actor_id=actor_id,
|
|
466
|
+
reason=reason,
|
|
467
|
+
delete_message_seconds=delete_message_seconds,
|
|
468
|
+
)
|
|
469
|
+
|
|
470
|
+
async def kick(self, guild_id: int, user_id: int, *, actor_id: int | None, reason: str) -> None:
|
|
471
|
+
self._record("kick", guild_id, user_id, actor_id=actor_id, reason=reason)
|
|
472
|
+
|
|
473
|
+
async def timeout(
|
|
474
|
+
self,
|
|
475
|
+
guild_id: int,
|
|
476
|
+
user_id: int,
|
|
477
|
+
*,
|
|
478
|
+
actor_id: int | None,
|
|
479
|
+
reason: str,
|
|
480
|
+
duration_seconds: int,
|
|
481
|
+
) -> None:
|
|
482
|
+
self._record(
|
|
483
|
+
"timeout",
|
|
484
|
+
guild_id,
|
|
485
|
+
user_id,
|
|
486
|
+
actor_id=actor_id,
|
|
487
|
+
reason=reason,
|
|
488
|
+
duration_seconds=duration_seconds,
|
|
489
|
+
)
|
|
490
|
+
|
|
491
|
+
async def fetch_message(self, ref: MessageRef) -> MessageSnapshot | None:
|
|
492
|
+
self._record("fetch_message", ref)
|
|
493
|
+
return self.messages.get(ref)
|
|
494
|
+
|
|
495
|
+
async def fetch_member(self, guild_id: int, user_id: int) -> MemberSnapshot | None:
|
|
496
|
+
self._record("fetch_member", guild_id, user_id)
|
|
497
|
+
return self.members.get((guild_id, user_id))
|
|
498
|
+
|
|
499
|
+
async def fetch_channel(self, guild_id: int, channel_id: int) -> ChannelSnapshot | None:
|
|
500
|
+
self._record("fetch_channel", guild_id, channel_id)
|
|
501
|
+
return self.channels.get((guild_id, channel_id))
|
|
502
|
+
|
|
503
|
+
async def fetch_messages(
|
|
504
|
+
self,
|
|
505
|
+
guild_id: int,
|
|
506
|
+
channel_id: int,
|
|
507
|
+
*,
|
|
508
|
+
after_message_id: int | None = None,
|
|
509
|
+
before_message_id: int | None = None,
|
|
510
|
+
limit: int = 100,
|
|
511
|
+
) -> MessagePage:
|
|
512
|
+
self._record(
|
|
513
|
+
"fetch_messages",
|
|
514
|
+
guild_id,
|
|
515
|
+
channel_id,
|
|
516
|
+
after_message_id=after_message_id,
|
|
517
|
+
before_message_id=before_message_id,
|
|
518
|
+
limit=limit,
|
|
519
|
+
)
|
|
520
|
+
messages = sorted(
|
|
521
|
+
self.histories.get((guild_id, channel_id), ()),
|
|
522
|
+
key=lambda message: message.ref.message_id,
|
|
523
|
+
)
|
|
524
|
+
if after_message_id is not None:
|
|
525
|
+
candidates = [
|
|
526
|
+
message for message in messages if message.ref.message_id > after_message_id
|
|
527
|
+
]
|
|
528
|
+
selected = candidates[:limit]
|
|
529
|
+
cursor = selected[-1].ref.message_id if selected else None
|
|
530
|
+
else:
|
|
531
|
+
candidates = [
|
|
532
|
+
message
|
|
533
|
+
for message in messages
|
|
534
|
+
if before_message_id is None or message.ref.message_id < before_message_id
|
|
535
|
+
]
|
|
536
|
+
selected = candidates[-limit:]
|
|
537
|
+
cursor = selected[0].ref.message_id if selected else None
|
|
538
|
+
return MessagePage(tuple(selected), cursor, len(candidates) > len(selected))
|
|
539
|
+
|
|
540
|
+
async def fetch_pins(self, guild_id: int, channel_id: int) -> tuple[MessageSnapshot, ...]:
|
|
541
|
+
self._record("fetch_pins", guild_id, channel_id)
|
|
542
|
+
return self.pins.get((guild_id, channel_id), ())
|
|
543
|
+
|
|
544
|
+
async def fetch_public_threads(
|
|
545
|
+
self, guild_id: int, parent_channel_id: int
|
|
546
|
+
) -> tuple[ChannelSnapshot, ...]:
|
|
547
|
+
self._record("fetch_public_threads", guild_id, parent_channel_id)
|
|
548
|
+
return self.public_threads.get((guild_id, parent_channel_id), ())
|
|
549
|
+
|
|
550
|
+
async def fetch_roles(self, guild_id: int) -> tuple[RoleSnapshot, ...]:
|
|
551
|
+
self._record("fetch_roles", guild_id)
|
|
552
|
+
return self.roles.get(guild_id, ())
|
|
553
|
+
|
|
554
|
+
async def fetch_invites(self, guild_id: int) -> tuple[InviteSnapshot, ...]:
|
|
555
|
+
self._record("fetch_invites", guild_id)
|
|
556
|
+
return self.invites.get(guild_id, ())
|
|
557
|
+
|
|
558
|
+
async def can_view_channel(self, guild_id: int, user_id: int, channel_id: int) -> bool:
|
|
559
|
+
self._record("can_view_channel", guild_id, user_id, channel_id)
|
|
560
|
+
return self.channel_access.get((guild_id, user_id, channel_id), False)
|
|
561
|
+
|
|
562
|
+
|
|
563
|
+
@dataclass(slots=True)
|
|
564
|
+
class FakeResponse:
|
|
565
|
+
content: str | None
|
|
566
|
+
embed: OutgoingEmbed | None
|
|
567
|
+
ephemeral: bool
|
|
568
|
+
components: tuple[Any, ...]
|
|
569
|
+
kind: str
|
|
570
|
+
|
|
571
|
+
|
|
572
|
+
class FakeInteraction:
|
|
573
|
+
def __init__(
|
|
574
|
+
self,
|
|
575
|
+
*,
|
|
576
|
+
guild_id: int = 1,
|
|
577
|
+
channel_id: int = 2,
|
|
578
|
+
user_id: int = 3,
|
|
579
|
+
options: Mapping[str, Any] | None = None,
|
|
580
|
+
custom_id: str | None = None,
|
|
581
|
+
values: Sequence[str] = (),
|
|
582
|
+
guild_name: str | None = "Test Guild",
|
|
583
|
+
message: MessageRef | None = None,
|
|
584
|
+
) -> None:
|
|
585
|
+
self._message = message
|
|
586
|
+
self._guild_name = guild_name
|
|
587
|
+
self._guild_id = guild_id
|
|
588
|
+
self._channel_id = channel_id
|
|
589
|
+
self._user_id = user_id
|
|
590
|
+
self._options = dict(options or {})
|
|
591
|
+
self._custom_id = custom_id
|
|
592
|
+
self._values = tuple(values)
|
|
593
|
+
self.responses: list[FakeResponse] = []
|
|
594
|
+
self.deferred: bool | None = None
|
|
595
|
+
|
|
596
|
+
@property
|
|
597
|
+
def guild_id(self) -> int:
|
|
598
|
+
return self._guild_id
|
|
599
|
+
|
|
600
|
+
@property
|
|
601
|
+
def channel_id(self) -> int:
|
|
602
|
+
return self._channel_id
|
|
603
|
+
|
|
604
|
+
@property
|
|
605
|
+
def user_id(self) -> int:
|
|
606
|
+
return self._user_id
|
|
607
|
+
|
|
608
|
+
@property
|
|
609
|
+
def guild_name(self) -> str | None:
|
|
610
|
+
return self._guild_name
|
|
611
|
+
|
|
612
|
+
@property
|
|
613
|
+
def options(self) -> Mapping[str, Any]:
|
|
614
|
+
return self._options
|
|
615
|
+
|
|
616
|
+
@property
|
|
617
|
+
def custom_id(self) -> str | None:
|
|
618
|
+
return self._custom_id
|
|
619
|
+
|
|
620
|
+
@property
|
|
621
|
+
def values(self) -> tuple[str, ...]:
|
|
622
|
+
return self._values
|
|
623
|
+
|
|
624
|
+
@property
|
|
625
|
+
def message(self) -> MessageRef | None:
|
|
626
|
+
return self._message
|
|
627
|
+
|
|
628
|
+
async def respond(
|
|
629
|
+
self,
|
|
630
|
+
content: str | None = None,
|
|
631
|
+
*,
|
|
632
|
+
embed: OutgoingEmbed | None = None,
|
|
633
|
+
ephemeral: bool = False,
|
|
634
|
+
components: Sequence[Any] = (),
|
|
635
|
+
) -> None:
|
|
636
|
+
self.responses.append(FakeResponse(content, embed, ephemeral, tuple(components), "respond"))
|
|
637
|
+
|
|
638
|
+
async def defer(self, *, ephemeral: bool = False) -> None:
|
|
639
|
+
self.deferred = ephemeral
|
|
640
|
+
|
|
641
|
+
async def edit_original(
|
|
642
|
+
self,
|
|
643
|
+
content: str | None = None,
|
|
644
|
+
*,
|
|
645
|
+
embed: OutgoingEmbed | None = None,
|
|
646
|
+
components: Sequence[Any] = (),
|
|
647
|
+
) -> None:
|
|
648
|
+
self.responses.append(FakeResponse(content, embed, False, tuple(components), "edit"))
|
|
649
|
+
|
|
650
|
+
async def follow_up(
|
|
651
|
+
self, content: str, *, embed: OutgoingEmbed | None = None, ephemeral: bool = False
|
|
652
|
+
) -> None:
|
|
653
|
+
self.responses.append(FakeResponse(content, embed, ephemeral, (), "follow_up"))
|
|
654
|
+
|
|
655
|
+
@property
|
|
656
|
+
def last(self) -> FakeResponse:
|
|
657
|
+
return self.responses[-1]
|
|
658
|
+
|
|
659
|
+
|
|
660
|
+
class FakeInteractions:
|
|
661
|
+
"""Records command and component registrations; tests invoke handlers directly."""
|
|
662
|
+
|
|
663
|
+
def __init__(self, module_name: str) -> None:
|
|
664
|
+
self.module_name = module_name
|
|
665
|
+
self.commands: dict[str, tuple[CommandSpec, Callable[..., Any]]] = {}
|
|
666
|
+
self.components: dict[tuple[str, str], Callable[..., Any]] = {}
|
|
667
|
+
self.component_min_tiers: dict[tuple[str, str], TrustTierName] = {}
|
|
668
|
+
self.autocompletes: dict[str, Callable[..., Any]] = {}
|
|
669
|
+
|
|
670
|
+
def add_command(
|
|
671
|
+
self,
|
|
672
|
+
spec: CommandSpec,
|
|
673
|
+
handler: Callable[..., Any],
|
|
674
|
+
*,
|
|
675
|
+
autocomplete: Callable[..., Any] | None = None,
|
|
676
|
+
) -> _Closable:
|
|
677
|
+
qualified = f"{spec.group}.{spec.name}" if spec.group else spec.name
|
|
678
|
+
if qualified in self.commands:
|
|
679
|
+
raise ModuleContractError(
|
|
680
|
+
f"module {self.module_name!r} command {qualified!r} is already registered"
|
|
681
|
+
)
|
|
682
|
+
self.commands[qualified] = (spec, handler)
|
|
683
|
+
if autocomplete is not None:
|
|
684
|
+
self.autocompletes[qualified] = autocomplete
|
|
685
|
+
return _Closable(
|
|
686
|
+
lambda: (self.commands.pop(qualified, None), self.autocompletes.pop(qualified, None))
|
|
687
|
+
)
|
|
688
|
+
|
|
689
|
+
def register_component(
|
|
690
|
+
self,
|
|
691
|
+
kind: str,
|
|
692
|
+
key: str,
|
|
693
|
+
handler: Callable[..., Any],
|
|
694
|
+
*,
|
|
695
|
+
expires_after_seconds: float | None = None,
|
|
696
|
+
min_tier: TrustTierName = "member",
|
|
697
|
+
) -> _Closable:
|
|
698
|
+
if kind not in ("button", "select"):
|
|
699
|
+
raise ModuleContractError(f"unsupported component kind {kind!r}")
|
|
700
|
+
build_custom_id(self.module_name, key)
|
|
701
|
+
identity = (kind, key)
|
|
702
|
+
if identity in self.components:
|
|
703
|
+
raise ModuleContractError(
|
|
704
|
+
f"module {self.module_name!r} component {kind}/{key!r} is already registered"
|
|
705
|
+
)
|
|
706
|
+
self.components[identity] = handler
|
|
707
|
+
self.component_min_tiers[identity] = min_tier
|
|
708
|
+
return _Closable(
|
|
709
|
+
lambda: (
|
|
710
|
+
self.components.pop(identity, None),
|
|
711
|
+
self.component_min_tiers.pop(identity, None),
|
|
712
|
+
)
|
|
713
|
+
)
|
|
714
|
+
|
|
715
|
+
def custom_id(self, key: str, *parts: str) -> str:
|
|
716
|
+
return build_custom_id(self.module_name, key, *parts)
|
|
717
|
+
|
|
718
|
+
|
|
719
|
+
class FakeGuildSettings:
|
|
720
|
+
def __init__(
|
|
721
|
+
self, values: Mapping[int, Mapping[str, Any]] | None = None, *, enabled: bool = True
|
|
722
|
+
) -> None:
|
|
723
|
+
self.values: dict[int, dict[str, Any]] = {gid: dict(v) for gid, v in (values or {}).items()}
|
|
724
|
+
self.enabled = enabled
|
|
725
|
+
self.errors: dict[int, tuple[str, ...]] = {}
|
|
726
|
+
self._callbacks: list[Callable[[int], None]] = []
|
|
727
|
+
|
|
728
|
+
def guild_ids(self) -> tuple[int, ...]:
|
|
729
|
+
return tuple(sorted(self.values))
|
|
730
|
+
|
|
731
|
+
def get(self, guild_id: int) -> GuildSettingsSnapshot:
|
|
732
|
+
errors = self.errors.get(guild_id, ())
|
|
733
|
+
return GuildSettingsSnapshot(
|
|
734
|
+
self.values.get(guild_id, {}), not errors, errors, f"rev-{guild_id}"
|
|
735
|
+
)
|
|
736
|
+
|
|
737
|
+
def is_enabled(self, guild_id: int) -> bool:
|
|
738
|
+
return self.enabled and not self.errors.get(guild_id)
|
|
739
|
+
|
|
740
|
+
def on_change(self, callback: Callable[[int], None]) -> _Closable:
|
|
741
|
+
self._callbacks.append(callback)
|
|
742
|
+
return _Closable(lambda: self._callbacks.remove(callback))
|
|
743
|
+
|
|
744
|
+
def set(self, guild_id: int, **values: Any) -> None:
|
|
745
|
+
self.values.setdefault(guild_id, {}).update(values)
|
|
746
|
+
for callback in list(self._callbacks):
|
|
747
|
+
callback(guild_id)
|
|
748
|
+
|
|
749
|
+
|
|
750
|
+
type FakeRoute = Callable[[str, Mapping[str, str]], HttpResponse]
|
|
751
|
+
|
|
752
|
+
|
|
753
|
+
class FakeHttp:
|
|
754
|
+
"""Routes are matched by URL prefix; unmatched hosts raise ``HostNotAllowed``."""
|
|
755
|
+
|
|
756
|
+
def __init__(
|
|
757
|
+
self, routes: Mapping[str, HttpResponse | bytes | FakeRoute] | None = None
|
|
758
|
+
) -> None:
|
|
759
|
+
self.routes: dict[str, HttpResponse | bytes | FakeRoute] = dict(routes or {})
|
|
760
|
+
self.requests: list[tuple[str, str, Mapping[str, str]]] = []
|
|
761
|
+
|
|
762
|
+
def _resolve(self, method: str, url: str, headers: Mapping[str, str] | None) -> HttpResponse:
|
|
763
|
+
sent = dict(headers or {})
|
|
764
|
+
self.requests.append((method, url, sent))
|
|
765
|
+
for prefix, route in self.routes.items():
|
|
766
|
+
if url.startswith(prefix):
|
|
767
|
+
if callable(route):
|
|
768
|
+
return route(url, sent)
|
|
769
|
+
if isinstance(route, bytes):
|
|
770
|
+
return HttpResponse(200, {}, route)
|
|
771
|
+
return route
|
|
772
|
+
raise HostNotAllowed(f"no fake route for {url!r}")
|
|
773
|
+
|
|
774
|
+
async def get(
|
|
775
|
+
self,
|
|
776
|
+
url: str,
|
|
777
|
+
*,
|
|
778
|
+
headers: Mapping[str, str] | None = None,
|
|
779
|
+
timeout_seconds: float = 20.0,
|
|
780
|
+
max_bytes: int = 8 * 1024 * 1024,
|
|
781
|
+
) -> HttpResponse:
|
|
782
|
+
return self._resolve("GET", url, headers)
|
|
783
|
+
|
|
784
|
+
async def post_json(
|
|
785
|
+
self,
|
|
786
|
+
url: str,
|
|
787
|
+
payload: Any,
|
|
788
|
+
*,
|
|
789
|
+
headers: Mapping[str, str] | None = None,
|
|
790
|
+
timeout_seconds: float = 20.0,
|
|
791
|
+
max_bytes: int = 8 * 1024 * 1024,
|
|
792
|
+
) -> HttpResponse:
|
|
793
|
+
return self._resolve("POST", url, headers)
|
|
794
|
+
|
|
795
|
+
async def download(
|
|
796
|
+
self,
|
|
797
|
+
url: str,
|
|
798
|
+
*,
|
|
799
|
+
headers: Mapping[str, str] | None = None,
|
|
800
|
+
timeout_seconds: float = 30.0,
|
|
801
|
+
max_bytes: int = 8 * 1024 * 1024,
|
|
802
|
+
) -> AsyncIterator[bytes]:
|
|
803
|
+
response = self._resolve("GET", url, headers)
|
|
804
|
+
yield response.body
|
|
805
|
+
|
|
806
|
+
|
|
807
|
+
@dataclass(slots=True)
|
|
808
|
+
class _FakeProvided:
|
|
809
|
+
implementation: object
|
|
810
|
+
alive: bool = True
|
|
811
|
+
|
|
812
|
+
|
|
813
|
+
class _FakeServiceProxy:
|
|
814
|
+
"""Forwards attributes like the host proxy; dead for good once its registration closes."""
|
|
815
|
+
|
|
816
|
+
__slots__ = ("_key", "_provided")
|
|
817
|
+
|
|
818
|
+
def __init__(self, provided: _FakeProvided, key: tuple[str, int]) -> None:
|
|
819
|
+
self._provided = provided
|
|
820
|
+
self._key = key
|
|
821
|
+
|
|
822
|
+
def __getattr__(self, attribute: str) -> Any:
|
|
823
|
+
if not self._provided.alive:
|
|
824
|
+
name, version = self._key
|
|
825
|
+
raise ServiceUnavailable(f"service {name}@{version} closed")
|
|
826
|
+
return getattr(self._provided.implementation, attribute)
|
|
827
|
+
|
|
828
|
+
|
|
829
|
+
class FakeServiceRegistry:
|
|
830
|
+
def __init__(self) -> None:
|
|
831
|
+
self.provided: dict[tuple[str, int], object] = {}
|
|
832
|
+
self._records: dict[tuple[str, int], _FakeProvided] = {}
|
|
833
|
+
|
|
834
|
+
def provide(self, name: str, version: int, implementation: object) -> _Closable:
|
|
835
|
+
key = (name, version)
|
|
836
|
+
live = self._records.get(key)
|
|
837
|
+
if live is not None and live.alive:
|
|
838
|
+
# The host refuses a second live provider for the same service.
|
|
839
|
+
raise ModuleContractError(f"service {name}@{version} is already provided")
|
|
840
|
+
record = _FakeProvided(implementation)
|
|
841
|
+
self.provided[key] = implementation
|
|
842
|
+
self._records[key] = record
|
|
843
|
+
|
|
844
|
+
def close() -> None:
|
|
845
|
+
# Like the host, a re-provided service is a new object; proxies
|
|
846
|
+
# handed out for this one stay closed.
|
|
847
|
+
record.alive = False
|
|
848
|
+
if self._records.get(key) is record:
|
|
849
|
+
self._records.pop(key, None)
|
|
850
|
+
self.provided.pop(key, None)
|
|
851
|
+
|
|
852
|
+
return _Closable(close)
|
|
853
|
+
|
|
854
|
+
@overload
|
|
855
|
+
def get(self, name: str, version: int) -> object: ...
|
|
856
|
+
|
|
857
|
+
@overload
|
|
858
|
+
def get(self, name: str, version: int, type_: type[_T]) -> _T: ...
|
|
859
|
+
|
|
860
|
+
def get(self, name: str, version: int, type_: type[_T] | None = None) -> object:
|
|
861
|
+
try:
|
|
862
|
+
record = self._records[(name, version)]
|
|
863
|
+
except KeyError as exc:
|
|
864
|
+
raise ServiceUnavailable(f"service {name}@{version} is not provided") from exc
|
|
865
|
+
if type_ is not None and not isinstance(record.implementation, type_):
|
|
866
|
+
raise TypeError(f"service {name}@{version} is not a {type_.__name__}")
|
|
867
|
+
# A proxy, as in the host: attribute access only, dead after close.
|
|
868
|
+
return _FakeServiceProxy(record, (name, version))
|
|
869
|
+
|
|
870
|
+
|
|
871
|
+
class FakeTrust:
|
|
872
|
+
def __init__(
|
|
873
|
+
self,
|
|
874
|
+
tiers: Mapping[tuple[int, int], TrustTierName] | None = None,
|
|
875
|
+
*,
|
|
876
|
+
default: TrustTierName = "member",
|
|
877
|
+
) -> None:
|
|
878
|
+
self.tiers: dict[tuple[int, int], TrustTierName] = dict(tiers or {})
|
|
879
|
+
self.default = default
|
|
880
|
+
|
|
881
|
+
async def tier(self, guild_id: int, user_id: int) -> TrustTierName:
|
|
882
|
+
return self.tiers.get((guild_id, user_id), self.default)
|
|
883
|
+
|
|
884
|
+
|
|
885
|
+
class MemoryStorage:
|
|
886
|
+
"""``ModuleStorage`` over one in-memory SQLite connection.
|
|
887
|
+
|
|
888
|
+
Requires the ``testing`` extra (``kimi-agent-module-api[testing]``), which
|
|
889
|
+
brings ``aiosqlite``. Mirrors the host's guarantees: ``table()`` returns
|
|
890
|
+
the quoted ``"<module>_<name>"``, reads go through ``connection``, and
|
|
891
|
+
``write_transaction()`` serializes writers, commits on success, and rolls
|
|
892
|
+
back on error. Use ``migrate()`` to apply a module's ``scoped_migrations``.
|
|
893
|
+
|
|
894
|
+
Typical use::
|
|
895
|
+
|
|
896
|
+
async with MemoryStorage.open("my_module") as storage:
|
|
897
|
+
await storage.migrate(MyModule.scoped_migrations)
|
|
898
|
+
...
|
|
899
|
+
"""
|
|
900
|
+
|
|
901
|
+
def __init__(self, connection: Any, module_name: str) -> None:
|
|
902
|
+
self._connection = connection
|
|
903
|
+
self._prefix = module_name.replace("-", "_")
|
|
904
|
+
self._write_lock = asyncio.Lock()
|
|
905
|
+
|
|
906
|
+
@classmethod
|
|
907
|
+
@contextlib.asynccontextmanager
|
|
908
|
+
async def open(cls, module_name: str) -> AsyncIterator[MemoryStorage]:
|
|
909
|
+
import aiosqlite # optional dependency: the ``testing`` extra
|
|
910
|
+
|
|
911
|
+
async with aiosqlite.connect(":memory:") as connection:
|
|
912
|
+
yield cls(connection, module_name)
|
|
913
|
+
|
|
914
|
+
@property
|
|
915
|
+
def connection(self) -> Any:
|
|
916
|
+
return self._connection
|
|
917
|
+
|
|
918
|
+
def table(self, name: str) -> str:
|
|
919
|
+
if not TABLE_NAME_RE.match(name):
|
|
920
|
+
raise ModuleContractError(f"table name {name!r} is not a valid identifier")
|
|
921
|
+
return f'"{self._prefix}_{name}"'
|
|
922
|
+
|
|
923
|
+
@contextlib.asynccontextmanager
|
|
924
|
+
async def _transaction(self) -> AsyncIterator[Any]:
|
|
925
|
+
async with self._write_lock:
|
|
926
|
+
try:
|
|
927
|
+
yield self._connection
|
|
928
|
+
except BaseException:
|
|
929
|
+
await self._connection.rollback()
|
|
930
|
+
raise
|
|
931
|
+
await self._connection.commit()
|
|
932
|
+
|
|
933
|
+
def write_transaction(self) -> Any:
|
|
934
|
+
return self._transaction()
|
|
935
|
+
|
|
936
|
+
async def migrate(self, migrations: Sequence[ScopedModuleMigration]) -> None:
|
|
937
|
+
for _name, migration in migrations:
|
|
938
|
+
await migration(MigrationContext(connection=self._connection, table=self.table))
|
|
939
|
+
await self._connection.commit()
|
|
940
|
+
|
|
941
|
+
|
|
942
|
+
@dataclass(slots=True)
|
|
943
|
+
class RecordedTool:
|
|
944
|
+
description: str
|
|
945
|
+
parameters: dict[str, Any]
|
|
946
|
+
handler: ModuleToolHandler
|
|
947
|
+
min_tier: TrustTier
|
|
948
|
+
searchable: bool
|
|
949
|
+
owner_only: bool
|
|
950
|
+
guild_only: bool
|
|
951
|
+
guild_ids: frozenset[int] | None
|
|
952
|
+
|
|
953
|
+
|
|
954
|
+
class RecordingToolRegistry:
|
|
955
|
+
"""Satisfies ``ModuleToolRegistry`` and remembers every registration."""
|
|
956
|
+
|
|
957
|
+
def __init__(self) -> None:
|
|
958
|
+
self.tools: dict[str, RecordedTool] = {}
|
|
959
|
+
|
|
960
|
+
def register(
|
|
961
|
+
self,
|
|
962
|
+
name: str,
|
|
963
|
+
description: str,
|
|
964
|
+
parameters: dict[str, Any],
|
|
965
|
+
handler: ModuleToolHandler,
|
|
966
|
+
*,
|
|
967
|
+
min_tier: TrustTier = TrustTier.MEMBER,
|
|
968
|
+
searchable: bool = False,
|
|
969
|
+
owner_only: bool = False,
|
|
970
|
+
guild_only: bool = True,
|
|
971
|
+
guild_ids: frozenset[int] | None = None,
|
|
972
|
+
) -> None:
|
|
973
|
+
self.tools[name] = RecordedTool(
|
|
974
|
+
description,
|
|
975
|
+
parameters,
|
|
976
|
+
handler,
|
|
977
|
+
min_tier,
|
|
978
|
+
searchable,
|
|
979
|
+
owner_only,
|
|
980
|
+
guild_only,
|
|
981
|
+
guild_ids,
|
|
982
|
+
)
|
|
983
|
+
|
|
984
|
+
|
|
985
|
+
@dataclass(slots=True)
|
|
986
|
+
class LoadContextRecorder:
|
|
987
|
+
"""What a ``load_context`` captured from ``create()``."""
|
|
988
|
+
|
|
989
|
+
registry: RecordingToolRegistry
|
|
990
|
+
labels: dict[str, str]
|
|
991
|
+
surfaces: dict[str, tuple[str, ...]]
|
|
992
|
+
|
|
993
|
+
|
|
994
|
+
def load_context(
|
|
995
|
+
settings: BaseSettings | None,
|
|
996
|
+
*,
|
|
997
|
+
capabilities: ModuleCapabilities | None = None,
|
|
998
|
+
registry: RecordingToolRegistry | None = None,
|
|
999
|
+
) -> tuple[ModuleLoadContext, LoadContextRecorder]:
|
|
1000
|
+
"""Build a ``ModuleLoadContext`` for calling ``ModuleSpec.create`` in a test."""
|
|
1001
|
+
recorder = LoadContextRecorder(registry or RecordingToolRegistry(), {}, {})
|
|
1002
|
+
|
|
1003
|
+
def declare(surface: str, names: Sequence[str]) -> None:
|
|
1004
|
+
recorder.surfaces[surface] = tuple(names)
|
|
1005
|
+
|
|
1006
|
+
context = ModuleLoadContext(
|
|
1007
|
+
capabilities=capabilities or ModuleCapabilities(BASELINE_CAPABILITIES, False, False),
|
|
1008
|
+
registry=recorder.registry,
|
|
1009
|
+
module_settings=settings,
|
|
1010
|
+
label_sink=recorder.labels.update,
|
|
1011
|
+
surface_sink=declare,
|
|
1012
|
+
)
|
|
1013
|
+
return context, recorder
|
|
1014
|
+
|
|
1015
|
+
|
|
1016
|
+
__all__ = [
|
|
1017
|
+
"DiscordCall",
|
|
1018
|
+
"FakeDiscordActions",
|
|
1019
|
+
"FakeEvents",
|
|
1020
|
+
"FakeGuildSettings",
|
|
1021
|
+
"FakeHealth",
|
|
1022
|
+
"FakeHttp",
|
|
1023
|
+
"FakeInteraction",
|
|
1024
|
+
"FakeInteractions",
|
|
1025
|
+
"FakeProposals",
|
|
1026
|
+
"FakeResponse",
|
|
1027
|
+
"FakeScheduler",
|
|
1028
|
+
"FakeServiceRegistry",
|
|
1029
|
+
"FakeTrust",
|
|
1030
|
+
"LoadContextRecorder",
|
|
1031
|
+
"MemoryStorage",
|
|
1032
|
+
"ProposedChange",
|
|
1033
|
+
"RecordedTool",
|
|
1034
|
+
"RecordingToolRegistry",
|
|
1035
|
+
"load_context",
|
|
1036
|
+
]
|