Skip to content

Commit 03bac12

Browse files
committed
feat: add addressed Discord gateway routing
Deliver Discord gateway traffic per guild and register persona guild ownership in Redis before executing actions. This keeps the one-socket-per-token model while avoiding broad event fan-out in discovery.
1 parent f6c0738 commit 03bac12

4 files changed

Lines changed: 130 additions & 14 deletions

File tree

flexus_client_kit/ckit_automation_actions.py

Lines changed: 16 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -1,10 +1,14 @@
11
"""
2-
Execute resolved automation actions (Discord + Mongo) and build job handlers for scheduled rules.
2+
Execute resolved automation actions and build job handlers for scheduled rules.
33
44
The automation engine (ckit_automation_engine) produces flat action dicts with pre-resolved
55
_resolved_body / _resolved_channel_id. This module performs side effects only, returns per-action
66
results for logging, and emits field_changes for crm_field_changed / status_transition cascades.
7-
Aligned with unified Discord bot plan U2.3.
7+
8+
Generic CRM actions (set_crm_field, set_status, enqueue_check, cancel_pending_jobs) work with
9+
any connector. Discord-specific actions (send_dm, post_to_channel, add_role, remove_role, kick)
10+
delegate to the Discord connector via ctx['connector'].execute_action or the legacy direct-client
11+
path (ctx['discord_client'] / ctx['guild']) for backward compatibility.
812
"""
913

1014
from __future__ import annotations
@@ -84,7 +88,12 @@ async def _do_send_dm(action: dict, ctx: dict) -> Tuple[dict, Optional[dict]]:
8488
uid_s = str(member_doc.get("user_id", "") or "")
8589
if not uid_s:
8690
return (_result_dict(ok=False, error="missing_user_id"), None)
87-
result = await connector.execute_action("send_dm", {"user_id": uid_s, "text": body})
91+
dm_params: dict = {"user_id": uid_s, "text": body}
92+
sid = str(ctx.get("server_id") or "")
93+
if sid:
94+
# Propagate guild context so the gateway ACL can verify this persona's access.
95+
dm_params["server_id"] = sid
96+
result = await connector.execute_action("send_dm", dm_params)
8897
return (_result_dict(ok=result.ok, error=result.error), None)
8998
member_discord = ctx.get("member_discord")
9099
if member_discord is None:
@@ -502,11 +511,11 @@ async def _do_call_gatekeeper_tool(action: dict, ctx: dict) -> Tuple[dict, Optio
502511
if existing_app:
503512
app_id: Optional[str] = existing_app["application_id"]
504513
else:
505-
app_id = await ckit_person_domain.application_create_pending(
514+
app_id = await ckit_person_domain.application_create_pending(
506515
fclient,
507516
ws_id,
508517
person_id,
509-
source="discord_onboarding",
518+
source="discord_bot",
510519
platform="discord",
511520
payload={"guild_id": str(guild_id), "discord_user_id": discord_user_id},
512521
)
@@ -691,7 +700,7 @@ async def execute_actions(actions: List[dict], ctx: dict) -> Tuple[List[dict], L
691700
async def _run_cascade(
692701
*,
693702
db: Any,
694-
client: discord.Client | None,
703+
client: Any,
695704
persona_id: str,
696705
setup: dict,
697706
rules: List[dict],
@@ -782,7 +791,7 @@ def make_automation_job_handler(
782791
setup: dict,
783792
engine_process_fn: Callable[..., List[dict]],
784793
db: Any,
785-
client: discord.Client | None,
794+
client: Any,
786795
persona_id: str,
787796
disabled_rules_cache: Optional[DisabledRulesCache] = None,
788797
connector: Any = None,

flexus_client_kit/ckit_connector_discord_gateway.py

Lines changed: 34 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,8 +1,11 @@
11
from __future__ import annotations
22

3+
import asyncio
4+
import logging
35
from collections.abc import Awaitable, Callable, Iterable
46
from typing import Any
57

8+
from flexus_client_kit import ckit_shutdown
69
from flexus_client_kit.ckit_connector import (
710
ActionDescriptor,
811
ActionResult,
@@ -13,6 +16,8 @@
1316
from flexus_client_kit.ckit_connector_discord import DISCORD_ACTIONS, DISCORD_TRIGGERS
1417
from flexus_client_kit.gateway.ckit_gateway_redis import DiscordGatewayRedisSidecar
1518

19+
logger = logging.getLogger(__name__)
20+
1621

1722
class DiscordGatewayConnector(ChatConnector):
1823
def __init__(
@@ -28,6 +33,7 @@ def __init__(
2833
self._sidecar = sidecar or DiscordGatewayRedisSidecar(token)
2934
self._event_callback: Callable[[NormalizedEvent], Awaitable[None]] | None = None
3035
self._connected = False
36+
self._refresh_task: asyncio.Task[None] | None = None
3137

3238
@property
3339
def platform(self) -> str:
@@ -58,7 +64,12 @@ def format_mention(self, user_id: str) -> str:
5864
return "<@%s>" % (user_id,)
5965

6066
async def set_allowed_guild_ids(self, ids: Iterable[int]) -> None:
61-
self._allowed_guild_ids = {int(x) for x in ids}
67+
new_ids = {int(x) for x in ids}
68+
self._allowed_guild_ids = new_ids
69+
if self._connected:
70+
await self._sidecar.update_guild_channels(new_ids)
71+
if new_ids:
72+
await self._sidecar.register_persona_guilds(self._persona_id, new_ids)
6273

6374
async def update_guild_ids(self, ids: Iterable[int]) -> None:
6475
await self.set_allowed_guild_ids(ids)
@@ -96,11 +107,32 @@ async def _dispatch(self, event: NormalizedEvent) -> None:
96107
await cb(event)
97108

98109
async def connect(self) -> None:
99-
await self._sidecar.start_event_consumer(self._dispatch)
110+
await self._sidecar.start_event_consumer(self._dispatch, self._allowed_guild_ids)
111+
if self._allowed_guild_ids:
112+
await self._sidecar.register_persona_guilds(self._persona_id, self._allowed_guild_ids)
100113
self._connected = True
114+
self._refresh_task = asyncio.create_task(self._guild_refresh_loop())
115+
116+
async def _guild_refresh_loop(self) -> None:
117+
"""Re-register persona->guild TTL in Redis every 120s for long-lived processes."""
118+
while self._connected and not ckit_shutdown.shutdown_event.is_set():
119+
await ckit_shutdown.wait(120.0)
120+
if not self._connected:
121+
break
122+
if self._allowed_guild_ids:
123+
await self._sidecar.register_persona_guilds(self._persona_id, self._allowed_guild_ids)
101124

102125
async def disconnect(self) -> None:
103126
self._connected = False
127+
rt = self._refresh_task
128+
self._refresh_task = None
129+
if rt is not None:
130+
rt.cancel()
131+
try:
132+
await rt
133+
except asyncio.CancelledError:
134+
pass
135+
await self._sidecar.unregister_persona_guilds(self._persona_id)
104136
await self._sidecar.close()
105137

106138
async def execute_action(self, action_type: str, params: dict) -> ActionResult:

flexus_client_kit/gateway/ckit_gateway_redis.py

Lines changed: 68 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -25,11 +25,13 @@
2525
action_result_to_dict,
2626
channel_cmd_discord,
2727
channel_events_discord,
28+
channel_events_discord_guild,
2829
channel_reply_discord,
2930
gateway_instance_key_from_token,
3031
gateway_result_envelope_from_dict,
3132
normalized_event_from_dict,
3233
parse_event_envelope,
34+
registry_key_persona_guilds,
3335
)
3436

3537
logger = logging.getLogger(__name__)
@@ -78,6 +80,7 @@ def __init__(
7880
) -> None:
7981
self._token = token
8082
self._key = gateway_instance_key_from_token(token)
83+
# Legacy broadcast channel kept for direct-socket / dev-fallback reference.
8184
self._events_ch = channel_events_discord(self._key)
8285
self._cmd_ch = channel_cmd_discord(self._key)
8386
self._redis_cmd = redis_cmd
@@ -87,25 +90,83 @@ def __init__(
8790
self._stop = asyncio.Event()
8891
self._reader_task: asyncio.Task[None] | None = None
8992
self._cb: Callable[[NormalizedEvent], Awaitable[None]] | None = None
93+
# Live pubsub handle shared between reader loop and update_guild_channels.
94+
self._ps: Any | None = None
95+
# Currently active per-guild subscription channels.
96+
self._subscribed_guild_channels: set[str] = set()
9097

9198
@property
9299
def gateway_instance_key(self) -> str:
93100
return self._key
94101

95102
@property
96103
def events_channel(self) -> str:
104+
# Returns legacy broadcast channel for reference; production path uses per-guild channels.
97105
return self._events_ch
98106

99107
@property
100108
def cmd_channel(self) -> str:
101109
return self._cmd_ch
102110

103-
async def start_event_consumer(self, on_event: Callable[[NormalizedEvent], Awaitable[None]]) -> None:
111+
async def start_event_consumer(
112+
self,
113+
on_event: Callable[[NormalizedEvent], Awaitable[None]],
114+
guild_ids: set[int] | None = None,
115+
) -> None:
104116
self._cb = on_event
105117
self._stop.clear()
106118
if self._redis_pubsub is None:
107119
self._redis_pubsub = redis_pubsub_client_from_env()
108-
self._reader_task = asyncio.create_task(self._event_reader_loop())
120+
initial_channels = {
121+
channel_events_discord_guild(self._key, str(g))
122+
for g in (guild_ids or set())
123+
}
124+
self._subscribed_guild_channels = initial_channels
125+
self._reader_task = asyncio.create_task(self._event_reader_loop(initial_channels))
126+
127+
async def update_guild_channels(self, guild_ids: set[int]) -> None:
128+
"""Subscribe to new per-guild channels and drop channels no longer in the allowed set."""
129+
new_channels = {channel_events_discord_guild(self._key, str(g)) for g in guild_ids}
130+
to_add = new_channels - self._subscribed_guild_channels
131+
to_remove = self._subscribed_guild_channels - new_channels
132+
ps = self._ps
133+
if ps is None:
134+
# Reader loop not yet running; track for when it starts.
135+
self._subscribed_guild_channels = new_channels
136+
return
137+
try:
138+
if to_add:
139+
await ps.subscribe(*to_add)
140+
if to_remove:
141+
await ps.unsubscribe(*to_remove)
142+
except (RedisConnectionError, OSError, RuntimeError) as e:
143+
logger.warning("update_guild_channels: %s %s", type(e).__name__, e)
144+
self._subscribed_guild_channels = new_channels
145+
146+
async def register_persona_guilds(self, persona_id: str, guild_ids: set[int]) -> None:
147+
"""Write persona->guild Set in Redis with a 300s TTL; called on connect and refresh."""
148+
if self._redis_cmd is None:
149+
self._redis_cmd = redis_client_from_env()
150+
reg_key = registry_key_persona_guilds(self._key, persona_id)
151+
try:
152+
pipe = self._redis_cmd.pipeline()
153+
pipe.delete(reg_key)
154+
if guild_ids:
155+
pipe.sadd(reg_key, *[str(g) for g in guild_ids])
156+
pipe.expire(reg_key, 300)
157+
await pipe.execute()
158+
except (RedisConnectionError, RedisTimeoutError, OSError, RuntimeError) as e:
159+
logger.warning("register_persona_guilds: %s %s", type(e).__name__, e)
160+
161+
async def unregister_persona_guilds(self, persona_id: str) -> None:
162+
"""Remove persona->guild Set from Redis on clean disconnect."""
163+
if self._redis_cmd is None:
164+
return
165+
reg_key = registry_key_persona_guilds(self._key, persona_id)
166+
try:
167+
await self._redis_cmd.delete(reg_key)
168+
except (RedisConnectionError, RedisTimeoutError, OSError, RuntimeError) as e:
169+
logger.warning("unregister_persona_guilds: %s %s", type(e).__name__, e)
109170

110171
async def stop_event_consumer(self) -> None:
111172
self._stop.set()
@@ -127,12 +188,14 @@ async def close(self) -> None:
127188
await self._redis_cmd.close()
128189
self._redis_cmd = None
129190

130-
async def _event_reader_loop(self) -> None:
191+
async def _event_reader_loop(self, initial_channels: set[str]) -> None:
131192
r = self._redis_pubsub
132193
if r is None:
133194
return
134195
ps = r.pubsub()
135-
await ps.subscribe(self._events_ch)
196+
self._ps = ps
197+
if initial_channels:
198+
await ps.subscribe(*initial_channels)
136199
try:
137200
while not self._stop.is_set() and not ckit_shutdown.shutdown_event.is_set():
138201
try:
@@ -161,8 +224,8 @@ async def _event_reader_loop(self) -> None:
161224
except asyncio.CancelledError:
162225
raise
163226
finally:
227+
self._ps = None
164228
try:
165-
await ps.unsubscribe(self._events_ch)
166229
await ps.close()
167230
except (RedisConnectionError, OSError, RuntimeError):
168231
pass

flexus_client_kit/gateway/ckit_gateway_wire.py

Lines changed: 12 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -16,9 +16,21 @@ def gateway_instance_key_from_token(token: str) -> str:
1616

1717

1818
def channel_events_discord(gateway_instance_key: str) -> str:
19+
# Legacy broadcast channel — kept for direct-socket / dev-fallback mode only.
20+
# Production path uses channel_events_discord_guild for addressed per-guild delivery.
1921
return "gw:discord:%s:events" % (gateway_instance_key,)
2022

2123

24+
def channel_events_discord_guild(gateway_instance_key: str, guild_id: str) -> str:
25+
# Per-guild addressed event channel. Workers subscribe only to their allowed guilds.
26+
return "gw:discord:%s:guild:%s:events" % (gateway_instance_key, guild_id)
27+
28+
29+
def registry_key_persona_guilds(gateway_instance_key: str, persona_id: str) -> str:
30+
# Redis Set of guild_id strings for a persona; refreshed by the worker with a 300s TTL.
31+
return "gw:discord:%s:persona:%s:guilds" % (gateway_instance_key, persona_id)
32+
33+
2234
def channel_cmd_discord(gateway_instance_key: str) -> str:
2335
return "gw:discord:%s:cmd" % (gateway_instance_key,)
2436

0 commit comments

Comments
 (0)