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.
@@ -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
+ ]