Add Discord sharding support and related configurations
Enhance the bot's architecture by introducing sharding capabilities, allowing for better scalability and performance. Update environment files to include sharding options, and modify the bot's core logic to handle shard connections and latencies. Implement shard-specific logging and API synchronization, ensuring that only the primary shard performs certain tasks. Update API schemas and routes to reflect shard information, improving the overall status reporting and monitoring of the bot's performance across multiple shards.
This commit is contained in:
@@ -7,6 +7,12 @@ brand_name=Axiom
|
|||||||
OWNER_IDS=YOUR_DISCORD_USER_ID_HERE
|
OWNER_IDS=YOUR_DISCORD_USER_ID_HERE
|
||||||
DASHBOARD_ADMIN_IDS=YOUR_DISCORD_USER_ID_HERE
|
DASHBOARD_ADMIN_IDS=YOUR_DISCORD_USER_ID_HERE
|
||||||
|
|
||||||
|
# Sharding — leave unset for auto (recommended). Multi-process: same TOKEN, shared bot-db volume.
|
||||||
|
# SHARD_COUNT=
|
||||||
|
# SHARD_IDS=0,1
|
||||||
|
# SHARD_PRIMARY=
|
||||||
|
# Enable API only on the primary shard process.
|
||||||
|
|
||||||
# Shared secret — must match between bot API and dashboard
|
# Shared secret — must match between bot API and dashboard
|
||||||
DASHBOARD_API_KEY=generate_a_long_random_secret_here
|
DASHBOARD_API_KEY=generate_a_long_random_secret_here
|
||||||
|
|
||||||
|
|||||||
@@ -9,6 +9,16 @@ NEXT_PUBLIC_BRAND_NAME='Axiom'
|
|||||||
# BOT_INSTANCE_ID=prod-a
|
# BOT_INSTANCE_ID=prod-a
|
||||||
# API_PORT=8001
|
# API_PORT=8001
|
||||||
|
|
||||||
|
# ── Discord sharding (optional) ───────────────────────────────────────────────
|
||||||
|
# Default (all unset): AutoShardedBot asks Discord for recommended shard count (one process).
|
||||||
|
# Multi-process cluster: same TOKEN + shared DB volume; set SHARD_COUNT + SHARD_IDS per process.
|
||||||
|
# API / slash sync / emoji sync only on the primary (shard 0 in SHARD_IDS, or SHARD_PRIMARY=true).
|
||||||
|
# SHARD_COUNT=
|
||||||
|
# SHARD_IDS=0,1
|
||||||
|
# SHARD_PRIMARY=
|
||||||
|
# Only enable API on the primary process:
|
||||||
|
# API_ENABLED=true
|
||||||
|
|
||||||
# ── Owner / Staff IDs (REQUIRED — comma-separated Discord user IDs) ───────────
|
# ── Owner / Staff IDs (REQUIRED — comma-separated Discord user IDs) ───────────
|
||||||
# Set your own Discord user ID(s). No foreign IDs are used as fallback.
|
# Set your own Discord user ID(s). No foreign IDs are used as fallback.
|
||||||
OWNER_IDS=YOUR_DISCORD_USER_ID_HERE
|
OWNER_IDS=YOUR_DISCORD_USER_ID_HERE
|
||||||
|
|||||||
38
bot/Axiom.py
38
bot/Axiom.py
@@ -96,6 +96,17 @@ async def on_ready():
|
|||||||
print(f"Logged in as: {client.user}")
|
print(f"Logged in as: {client.user}")
|
||||||
print(f"Connected to: {len(client.guilds)} guilds")
|
print(f"Connected to: {len(client.guilds)} guilds")
|
||||||
print(f"Connected to: {len(client.users)} users")
|
print(f"Connected to: {len(client.users)} users")
|
||||||
|
shard_count = client.shard_count or 1
|
||||||
|
shard_ids = list(client.shards.keys()) if client.shards else [0]
|
||||||
|
print(f"Shards: {shard_count} (this process: {shard_ids})")
|
||||||
|
print(f"Shard primary: {IS_SHARD_PRIMARY}")
|
||||||
|
if client.latencies:
|
||||||
|
for sid, lat in client.latencies:
|
||||||
|
print(f" Shard {sid}: {round(lat * 1000)}ms")
|
||||||
|
|
||||||
|
if not IS_SHARD_PRIMARY:
|
||||||
|
print("\033[33m◈ Non-primary shard process — skipping API-side sync tasks\033[0m")
|
||||||
|
return
|
||||||
|
|
||||||
# Sync application emojis on startup
|
# Sync application emojis on startup
|
||||||
await run_sync(TOKEN)
|
await run_sync(TOKEN)
|
||||||
@@ -112,10 +123,31 @@ async def on_ready():
|
|||||||
client.loop.create_task(update_stats())
|
client.loop.create_task(update_stats())
|
||||||
|
|
||||||
|
|
||||||
|
@client.event
|
||||||
|
async def on_shard_connect(shard_id: int):
|
||||||
|
print(f"\033[32m◈ Shard {shard_id}: connected\033[0m")
|
||||||
|
|
||||||
|
|
||||||
|
@client.event
|
||||||
|
async def on_shard_ready(shard_id: int):
|
||||||
|
print(f"\033[32m◈ Shard {shard_id}: ready\033[0m")
|
||||||
|
|
||||||
|
|
||||||
|
@client.event
|
||||||
|
async def on_shard_disconnect(shard_id: int):
|
||||||
|
print(f"\033[33m◈ Shard {shard_id}: disconnected\033[0m")
|
||||||
|
|
||||||
|
|
||||||
|
@client.event
|
||||||
|
async def on_shard_resumed(shard_id: int):
|
||||||
|
print(f"\033[32m◈ Shard {shard_id}: resumed\033[0m")
|
||||||
|
|
||||||
|
|
||||||
@client.event
|
@client.event
|
||||||
async def on_guild_join(guild: discord.Guild):
|
async def on_guild_join(guild: discord.Guild):
|
||||||
# Only guild-mode needs per-server registration; global mode must not
|
# Only guild-mode needs per-server registration; global mode must not
|
||||||
# copy commands into the guild or Discord shows duplicates (/afk twice).
|
# copy commands into the guild or Discord shows duplicates (/afk twice).
|
||||||
|
if IS_SHARD_PRIMARY:
|
||||||
try:
|
try:
|
||||||
from utils.slash_sync import _sync_mode, sync_guild
|
from utils.slash_sync import _sync_mode, sync_guild
|
||||||
if _sync_mode() == "guild":
|
if _sync_mode() == "guild":
|
||||||
@@ -328,6 +360,9 @@ def run_api():
|
|||||||
uvicorn.run(fastapi_app, host=API_HOST, port=API_PORT, log_level="warning")
|
uvicorn.run(fastapi_app, host=API_HOST, port=API_PORT, log_level="warning")
|
||||||
|
|
||||||
def keep_alive():
|
def keep_alive():
|
||||||
|
if not IS_SHARD_PRIMARY:
|
||||||
|
print(f"\033[33m◈ API Server: Skipped (non-primary shard process)\033[0m")
|
||||||
|
return
|
||||||
if not API_ENABLED:
|
if not API_ENABLED:
|
||||||
print(f"\033[33m◈ API Server: Disabled via API_ENABLED=false\033[0m")
|
print(f"\033[33m◈ API Server: Disabled via API_ENABLED=false\033[0m")
|
||||||
return
|
return
|
||||||
@@ -341,7 +376,10 @@ keep_alive()
|
|||||||
|
|
||||||
# --- Cloudflare Tunnel (HTTPS for API) — only when explicitly enabled ---
|
# --- Cloudflare Tunnel (HTTPS for API) — only when explicitly enabled ---
|
||||||
from utils.tunnel import start_tunnel
|
from utils.tunnel import start_tunnel
|
||||||
|
if IS_SHARD_PRIMARY:
|
||||||
start_tunnel()
|
start_tunnel()
|
||||||
|
else:
|
||||||
|
print("\033[33m◈ Cloudflare Tunnel: Skipped (non-primary shard process)\033[0m")
|
||||||
|
|
||||||
# --- Main Bot Execution ---
|
# --- Main Bot Execution ---
|
||||||
async def main():
|
async def main():
|
||||||
|
|||||||
@@ -87,7 +87,16 @@ async def get_admin_stats(bot: "Axiom" = Depends(get_bot)):
|
|||||||
AdminNodeStatus(
|
AdminNodeStatus(
|
||||||
name="Auth Sockets",
|
name="Auth Sockets",
|
||||||
status="Healthy",
|
status="Healthy",
|
||||||
load=f"Shard: {bot.shard_count} | Latency: {round(bot.latency * 1000)}ms",
|
load=(
|
||||||
|
f"Shards: {bot.shard_count or 1} | "
|
||||||
|
f"Latency: {round(bot.latency * 1000)}ms"
|
||||||
|
+ (
|
||||||
|
" | "
|
||||||
|
+ ", ".join(f"s{sid}:{round(lat * 1000)}ms" for sid, lat in bot.latencies)
|
||||||
|
if bot.latencies and len(bot.latencies) > 1
|
||||||
|
else ""
|
||||||
|
)
|
||||||
|
),
|
||||||
icon="Lock"
|
icon="Lock"
|
||||||
)
|
)
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -28,13 +28,20 @@ async def get_status(bot: "Axiom" = Depends(get_bot)):
|
|||||||
"""
|
"""
|
||||||
Returns the live status of the bot.
|
Returns the live status of the bot.
|
||||||
"""
|
"""
|
||||||
|
shard_ids = list(bot.shards.keys()) if bot.shards else ([0] if bot.shard_count else [])
|
||||||
|
shard_latencies = [
|
||||||
|
{"id": sid, "latency_ms": round(lat * 1000, 2)}
|
||||||
|
for sid, lat in (bot.latencies or [])
|
||||||
|
]
|
||||||
return BotStatus(
|
return BotStatus(
|
||||||
user=str(bot.user),
|
user=str(bot.user),
|
||||||
id=str(bot.user.id) if bot.user else None,
|
id=str(bot.user.id) if bot.user else None,
|
||||||
latency=bot.latency * 1000,
|
latency=bot.latency * 1000,
|
||||||
guild_count=len(bot.guilds),
|
guild_count=len(bot.guilds),
|
||||||
user_count=sum(g.member_count or 0 for g in bot.guilds),
|
user_count=sum(g.member_count or 0 for g in bot.guilds),
|
||||||
shards=bot.shard_count
|
shards=bot.shard_count,
|
||||||
|
shard_ids=shard_ids,
|
||||||
|
shard_latencies=shard_latencies,
|
||||||
)
|
)
|
||||||
|
|
||||||
@router.get("/info", response_model=BotInfo, summary="Get bot info", description="Returns general information about the bot including command count and user reach.")
|
@router.get("/info", response_model=BotInfo, summary="Get bot info", description="Returns general information about the bot including command count and user reach.")
|
||||||
|
|||||||
@@ -31,6 +31,8 @@ class BotStatus(BaseModel):
|
|||||||
guild_count: int
|
guild_count: int
|
||||||
user_count: int
|
user_count: int
|
||||||
shards: Optional[int]
|
shards: Optional[int]
|
||||||
|
shard_ids: Optional[List[int]] = None
|
||||||
|
shard_latencies: Optional[List[Dict[str, Any]]] = None
|
||||||
|
|
||||||
# --- Guild Schemas ---
|
# --- Guild Schemas ---
|
||||||
|
|
||||||
|
|||||||
@@ -905,6 +905,11 @@ class Extra(commands.Cog):
|
|||||||
f"**API (Roundtrip):** `{api_latency}ms`\n"
|
f"**API (Roundtrip):** `{api_latency}ms`\n"
|
||||||
f"**Database:** `{db_latency}`"
|
f"**Database:** `{db_latency}`"
|
||||||
)
|
)
|
||||||
|
if self.bot.latencies and len(self.bot.latencies) > 1:
|
||||||
|
shard_lines = "\n".join(
|
||||||
|
f"**Shard {sid}:** `{round(lat * 1000)}ms`" for sid, lat in self.bot.latencies
|
||||||
|
)
|
||||||
|
latency_text += f"\n\n**Shards ({self.bot.shard_count or len(self.bot.latencies)}):**\n{shard_lines}"
|
||||||
await msg.edit(view=CV2("System Latency Report", latency_text))
|
await msg.edit(view=CV2("System Latency Report", latency_text))
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -19,7 +19,7 @@ import aiohttp
|
|||||||
import typing
|
import typing
|
||||||
from typing import List
|
from typing import List
|
||||||
import aiosqlite
|
import aiosqlite
|
||||||
from utils.config import OWNER_IDS, BotName
|
from utils.config import OWNER_IDS, BotName, get_shard_kwargs
|
||||||
from utils import getConfig, updateConfig
|
from utils import getConfig, updateConfig
|
||||||
from utils.Tools import setup_db
|
from utils.Tools import setup_db
|
||||||
from .Context import Context
|
from .Context import Context
|
||||||
@@ -39,6 +39,11 @@ class Axiom(commands.AutoShardedBot):
|
|||||||
intents = discord.Intents.all()
|
intents = discord.Intents.all()
|
||||||
intents.presences = True
|
intents.presences = True
|
||||||
intents.members = True
|
intents.members = True
|
||||||
|
shard_kwargs = get_shard_kwargs()
|
||||||
|
# Allow callers to override shard settings via kwargs
|
||||||
|
for key in ("shard_count", "shard_ids"):
|
||||||
|
if key in kwargs:
|
||||||
|
shard_kwargs[key] = kwargs.pop(key)
|
||||||
super().__init__(command_prefix=self.get_prefix,
|
super().__init__(command_prefix=self.get_prefix,
|
||||||
case_insensitive=True,
|
case_insensitive=True,
|
||||||
intents=intents,
|
intents=intents,
|
||||||
@@ -48,7 +53,7 @@ class Axiom(commands.AutoShardedBot):
|
|||||||
owner_ids=OWNER_IDS,
|
owner_ids=OWNER_IDS,
|
||||||
allowed_mentions=discord.AllowedMentions(
|
allowed_mentions=discord.AllowedMentions(
|
||||||
everyone=False, replied_user=False, roles=False),
|
everyone=False, replied_user=False, roles=False),
|
||||||
shard_count=1)
|
**shard_kwargs)
|
||||||
self.status_index = 0
|
self.status_index = 0
|
||||||
self.status_list = []
|
self.status_list = []
|
||||||
self._slash_synced = False
|
self._slash_synced = False
|
||||||
|
|||||||
@@ -80,6 +80,48 @@ OWNER_IDS_STR: list[str] = [str(i) for i in OWNER_IDS]
|
|||||||
# Discord application / OAuth client id (same as dashboard DISCORD_CLIENT_ID)
|
# Discord application / OAuth client id (same as dashboard DISCORD_CLIENT_ID)
|
||||||
DISCORD_CLIENT_ID = os.getenv("DISCORD_CLIENT_ID", "").strip()
|
DISCORD_CLIENT_ID = os.getenv("DISCORD_CLIENT_ID", "").strip()
|
||||||
|
|
||||||
|
# ── Discord sharding ──────────────────────────────────────────────────────────
|
||||||
|
# SHARD_COUNT empty = AutoShardedBot asks Discord for recommended count (in-process).
|
||||||
|
# SHARD_IDS (e.g. "0,1") + SHARD_COUNT = this process only runs those shards (cluster).
|
||||||
|
# Primary process runs API / slash sync / emoji sync / stats (shard 0 or override).
|
||||||
|
|
||||||
|
def _parse_optional_int(env_key: str) -> int | None:
|
||||||
|
raw = os.getenv(env_key, "").strip()
|
||||||
|
if not raw or not raw.isdigit():
|
||||||
|
return None
|
||||||
|
return int(raw)
|
||||||
|
|
||||||
|
|
||||||
|
SHARD_COUNT: int | None = _parse_optional_int("SHARD_COUNT")
|
||||||
|
SHARD_IDS: list[int] = _parse_ids("SHARD_IDS")
|
||||||
|
|
||||||
|
|
||||||
|
def get_shard_kwargs() -> dict:
|
||||||
|
"""Kwargs for AutoShardedBot.__init__ (omit keys for full auto-sharding)."""
|
||||||
|
kwargs: dict = {}
|
||||||
|
if SHARD_COUNT is not None:
|
||||||
|
kwargs["shard_count"] = SHARD_COUNT
|
||||||
|
if SHARD_IDS:
|
||||||
|
if SHARD_COUNT is None:
|
||||||
|
raise ValueError("SHARD_IDS requires SHARD_COUNT to be set")
|
||||||
|
kwargs["shard_ids"] = SHARD_IDS
|
||||||
|
kwargs["shard_count"] = SHARD_COUNT
|
||||||
|
return kwargs
|
||||||
|
|
||||||
|
|
||||||
|
def _resolve_is_shard_primary() -> bool:
|
||||||
|
raw = os.getenv("SHARD_PRIMARY", "").strip().lower()
|
||||||
|
if raw in ("1", "true", "yes", "on"):
|
||||||
|
return True
|
||||||
|
if raw in ("0", "false", "no", "off"):
|
||||||
|
return False
|
||||||
|
if not SHARD_IDS:
|
||||||
|
return True
|
||||||
|
return 0 in SHARD_IDS
|
||||||
|
|
||||||
|
|
||||||
|
IS_SHARD_PRIMARY: bool = _resolve_is_shard_primary()
|
||||||
|
|
||||||
# Aliases kept for backwards compatibility with files that import these names
|
# Aliases kept for backwards compatibility with files that import these names
|
||||||
BOT_OWNER_IDS = OWNER_IDS
|
BOT_OWNER_IDS = OWNER_IDS
|
||||||
BOT_OWNER_IDS_STR = OWNER_IDS_STR
|
BOT_OWNER_IDS_STR = OWNER_IDS_STR
|
||||||
|
|||||||
Reference in New Issue
Block a user