Initial commit
This commit is contained in:
0
bot/__init__.py
Normal file
0
bot/__init__.py
Normal file
62
bot/backlog.py
Normal file
62
bot/backlog.py
Normal file
@@ -0,0 +1,62 @@
|
||||
import asyncio
|
||||
import datetime
|
||||
import logging
|
||||
from typing import Awaitable, Callable
|
||||
|
||||
import discord
|
||||
|
||||
from bot.discovery import Target
|
||||
from bot.storage import Storage
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
ReactFn = Callable[[discord.Message, Target], Awaitable[bool]]
|
||||
|
||||
|
||||
async def scan_channel_backlog(
|
||||
storage: Storage,
|
||||
channel: discord.abc.Messageable,
|
||||
target: Target,
|
||||
friend_id: int,
|
||||
cutoff: datetime.datetime,
|
||||
react_fn: ReactFn,
|
||||
) -> None:
|
||||
"""Scan a channel's history once, resumably.
|
||||
|
||||
First run: walk forward from `cutoff` (now - days_back) to the present.
|
||||
Subsequent runs: only diff forward from the newest message already seen,
|
||||
since the days_back window only ever moves forward in time.
|
||||
"""
|
||||
channel_id = getattr(channel, "id", None)
|
||||
if channel_id is None:
|
||||
return
|
||||
|
||||
state = await asyncio.to_thread(storage.get_scan_state, channel_id)
|
||||
newest_id_seen = int(state["newest_id_seen"]) if state and state.get("newest_id_seen") else None
|
||||
|
||||
after = discord.Object(id=newest_id_seen) if newest_id_seen else cutoff
|
||||
|
||||
checkpoint_every = 25
|
||||
seen_since_checkpoint = 0
|
||||
|
||||
try:
|
||||
async for message in channel.history(after=after, oldest_first=True, limit=None):
|
||||
if newest_id_seen is None or message.id > newest_id_seen:
|
||||
newest_id_seen = message.id
|
||||
|
||||
await react_fn(message, target)
|
||||
|
||||
seen_since_checkpoint += 1
|
||||
if seen_since_checkpoint >= checkpoint_every:
|
||||
seen_since_checkpoint = 0
|
||||
await asyncio.to_thread(
|
||||
storage.set_scan_state, channel_id, target.target_id, newest_id_seen, False
|
||||
)
|
||||
|
||||
await asyncio.to_thread(
|
||||
storage.set_scan_state, channel_id, target.target_id, newest_id_seen, True
|
||||
)
|
||||
except discord.Forbidden:
|
||||
log.warning("No access to channel %s (%s), skipping", channel_id, getattr(channel, "name", ""))
|
||||
except discord.HTTPException:
|
||||
log.exception("HTTP error scanning channel %s", channel_id)
|
||||
14
bot/channels.py
Normal file
14
bot/channels.py
Normal file
@@ -0,0 +1,14 @@
|
||||
import logging
|
||||
from typing import List
|
||||
|
||||
import discord
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def get_scannable_channels(client: discord.Client, guild_id: int) -> List[discord.TextChannel]:
|
||||
guild = client.get_guild(guild_id)
|
||||
if guild is None:
|
||||
log.warning("Guild %s not found in local cache, skipping", guild_id)
|
||||
return []
|
||||
return list(guild.text_channels)
|
||||
40
bot/config.py
Normal file
40
bot/config.py
Normal file
@@ -0,0 +1,40 @@
|
||||
import json
|
||||
import os
|
||||
from dataclasses import dataclass, field
|
||||
from typing import List
|
||||
|
||||
|
||||
@dataclass
|
||||
class Config:
|
||||
token: str
|
||||
friend_id: int
|
||||
emoji: str
|
||||
days_back: int
|
||||
db_path: str
|
||||
log_level: str
|
||||
manual_guild_ids: List[int] = field(default_factory=list)
|
||||
|
||||
|
||||
def load_config(path: str = "config.json") -> Config:
|
||||
if not os.path.exists(path):
|
||||
raise FileNotFoundError(
|
||||
f"Config file not found: {path}. Copy config.example.json to config.json and fill it in."
|
||||
)
|
||||
|
||||
with open(path, "r", encoding="utf-8") as f:
|
||||
data = json.load(f)
|
||||
|
||||
required = ["token", "friend_id", "emoji"]
|
||||
missing = [key for key in required if not data.get(key)]
|
||||
if missing:
|
||||
raise ValueError(f"Missing required config fields: {missing}")
|
||||
|
||||
return Config(
|
||||
token=data["token"],
|
||||
friend_id=int(data["friend_id"]),
|
||||
emoji=data["emoji"],
|
||||
days_back=int(data.get("days_back", 10)),
|
||||
db_path=data.get("db_path", "reactions.sqlite3"),
|
||||
log_level=data.get("log_level", "INFO"),
|
||||
manual_guild_ids=[int(g) for g in data.get("manual_guild_ids", [])],
|
||||
)
|
||||
52
bot/discovery.py
Normal file
52
bot/discovery.py
Normal file
@@ -0,0 +1,52 @@
|
||||
import logging
|
||||
from dataclasses import dataclass
|
||||
from typing import List, Optional
|
||||
|
||||
import discord
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class Target:
|
||||
target_id: int
|
||||
target_type: str # "guild" | "group_dm"
|
||||
name: str
|
||||
|
||||
|
||||
async def discover_targets(client: discord.Client, friend_id: int, manual_guild_ids: Optional[List[int]] = None) -> List[Target]:
|
||||
targets: List[Target] = []
|
||||
guild_ids_seen = set()
|
||||
|
||||
try:
|
||||
friend = client.get_user(friend_id) or await client.fetch_user(friend_id)
|
||||
profile = await friend.profile(with_mutual_guilds=True, with_mutual_friends=False, with_mutual_friends_count=False)
|
||||
for mutual in (profile.mutual_guilds or []):
|
||||
guild = client.get_guild(mutual.id)
|
||||
name = guild.name if guild else f"Unknown Guild {mutual.id}"
|
||||
targets.append(Target(mutual.id, "guild", name))
|
||||
guild_ids_seen.add(mutual.id)
|
||||
except discord.NotFound:
|
||||
log.warning(
|
||||
"No mutual guild/friend relationship visible for user %s "
|
||||
"(profile lookup requires already being friends or sharing a server). "
|
||||
"Relying on manual_guild_ids and group DMs only.",
|
||||
friend_id,
|
||||
)
|
||||
except discord.HTTPException:
|
||||
log.warning("Mutual-guilds profile lookup failed; falling back to manual_guild_ids only", exc_info=True)
|
||||
|
||||
for guild_id in (manual_guild_ids or []):
|
||||
if guild_id in guild_ids_seen:
|
||||
continue
|
||||
guild = client.get_guild(guild_id)
|
||||
name = guild.name if guild else f"Manual Guild {guild_id}"
|
||||
targets.append(Target(guild_id, "guild", name))
|
||||
guild_ids_seen.add(guild_id)
|
||||
|
||||
for channel in client.private_channels:
|
||||
if isinstance(channel, discord.GroupChannel) and any(u.id == friend_id for u in channel.recipients):
|
||||
name = channel.name or ", ".join(u.name for u in channel.recipients)
|
||||
targets.append(Target(channel.id, "group_dm", name))
|
||||
|
||||
return targets
|
||||
62
bot/display.py
Normal file
62
bot/display.py
Normal file
@@ -0,0 +1,62 @@
|
||||
import asyncio
|
||||
from typing import List
|
||||
|
||||
from rich.console import Console
|
||||
from rich.live import Live
|
||||
from rich.panel import Panel
|
||||
from rich.table import Table
|
||||
|
||||
from bot.storage import Storage
|
||||
|
||||
|
||||
def _render(stats: List[dict]) -> Panel:
|
||||
table = Table(expand=True)
|
||||
table.add_column("Target")
|
||||
table.add_column("Type")
|
||||
table.add_column("Reactions", justify="right")
|
||||
table.add_column("Last reaction (UTC)")
|
||||
|
||||
total = 0
|
||||
for row in stats:
|
||||
total += row["reaction_count"]
|
||||
table.add_row(
|
||||
row["name"] or row["target_id"],
|
||||
"server" if row["target_type"] == "guild" else "group",
|
||||
str(row["reaction_count"]),
|
||||
row["last_reaction_at"] or "-",
|
||||
)
|
||||
|
||||
if not stats:
|
||||
table.add_row("(no targets discovered yet)", "-", "-", "-")
|
||||
|
||||
return Panel(table, title=f"Reaction stats — total: {total}", border_style="cyan")
|
||||
|
||||
|
||||
class StatsDisplay:
|
||||
"""A stats panel pinned to the bottom of the terminal, with normal
|
||||
log/print output scrolling above it (via rich.live.Live)."""
|
||||
|
||||
def __init__(self, console: Console, storage: Storage):
|
||||
self.console = console
|
||||
self.storage = storage
|
||||
self._live = Live(
|
||||
_render([]),
|
||||
console=console,
|
||||
refresh_per_second=4,
|
||||
transient=False,
|
||||
)
|
||||
|
||||
def __enter__(self):
|
||||
self._live.__enter__()
|
||||
return self
|
||||
|
||||
def __exit__(self, *exc_info):
|
||||
return self._live.__exit__(*exc_info)
|
||||
|
||||
def refresh_sync(self) -> None:
|
||||
stats = self.storage.get_stats()
|
||||
self._live.update(_render(stats))
|
||||
|
||||
async def refresh(self) -> None:
|
||||
stats = await asyncio.to_thread(self.storage.get_stats)
|
||||
self._live.update(_render(stats))
|
||||
63
bot/reactor.py
Normal file
63
bot/reactor.py
Normal file
@@ -0,0 +1,63 @@
|
||||
import asyncio
|
||||
import logging
|
||||
import random
|
||||
from typing import Optional
|
||||
|
||||
import discord
|
||||
|
||||
from bot.discovery import Target
|
||||
from bot.display import StatsDisplay
|
||||
from bot.storage import Storage
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def human_jitter() -> float:
|
||||
"""Gaussian delay centered a few seconds in, clamped to a natural-looking range."""
|
||||
delay = random.gauss(4.0, 1.5)
|
||||
return max(1.5, min(delay, 8.0))
|
||||
|
||||
|
||||
async def react_to_message(
|
||||
storage: Storage,
|
||||
message: discord.Message,
|
||||
target: Target,
|
||||
friend_id: int,
|
||||
emoji: str,
|
||||
source: str,
|
||||
display: Optional[StatsDisplay] = None,
|
||||
) -> bool:
|
||||
if message.author.id != friend_id:
|
||||
return False
|
||||
|
||||
already_reacted = any(
|
||||
reaction.me and str(reaction.emoji) == emoji for reaction in message.reactions
|
||||
)
|
||||
if not already_reacted:
|
||||
already_reacted = await asyncio.to_thread(storage.has_reacted, message.id, emoji)
|
||||
if already_reacted:
|
||||
return False
|
||||
|
||||
await asyncio.sleep(human_jitter())
|
||||
|
||||
try:
|
||||
await message.add_reaction(emoji)
|
||||
except discord.HTTPException:
|
||||
log.exception("Failed to react to message %s in channel %s", message.id, message.channel.id)
|
||||
return False
|
||||
|
||||
inserted = await asyncio.to_thread(
|
||||
storage.record_reaction,
|
||||
message.id,
|
||||
message.channel.id,
|
||||
target.target_id,
|
||||
friend_id,
|
||||
emoji,
|
||||
source,
|
||||
message.created_at,
|
||||
)
|
||||
if inserted:
|
||||
log.info("Reacted to message %s in %s [%s] (%s)", message.id, target.name, target.target_type, source)
|
||||
if display is not None:
|
||||
await display.refresh()
|
||||
return inserted
|
||||
157
bot/storage.py
Normal file
157
bot/storage.py
Normal file
@@ -0,0 +1,157 @@
|
||||
import sqlite3
|
||||
import threading
|
||||
from typing import Optional
|
||||
|
||||
|
||||
SCHEMA = """
|
||||
CREATE TABLE IF NOT EXISTS targets (
|
||||
target_id TEXT PRIMARY KEY,
|
||||
target_type TEXT NOT NULL CHECK(target_type IN ('guild','group_dm')),
|
||||
name TEXT,
|
||||
first_seen_at TEXT NOT NULL DEFAULT (strftime('%Y-%m-%dT%H:%M:%fZ','now')),
|
||||
reaction_count INTEGER NOT NULL DEFAULT 0
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS reactions_log (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
message_id TEXT NOT NULL,
|
||||
channel_id TEXT NOT NULL,
|
||||
target_id TEXT NOT NULL,
|
||||
author_id TEXT NOT NULL,
|
||||
emoji TEXT NOT NULL,
|
||||
reacted_at TEXT NOT NULL DEFAULT (strftime('%Y-%m-%dT%H:%M:%fZ','now')),
|
||||
message_created_at TEXT,
|
||||
source TEXT NOT NULL CHECK(source IN ('backlog','live')),
|
||||
UNIQUE(message_id, emoji)
|
||||
);
|
||||
CREATE INDEX IF NOT EXISTS idx_reactions_log_target ON reactions_log(target_id);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS channel_scan_state (
|
||||
channel_id TEXT PRIMARY KEY,
|
||||
target_id TEXT NOT NULL,
|
||||
newest_id_seen TEXT,
|
||||
backlog_complete INTEGER NOT NULL DEFAULT 0,
|
||||
last_scanned_at TEXT
|
||||
);
|
||||
"""
|
||||
|
||||
|
||||
class Storage:
|
||||
"""Thin synchronous sqlite3 wrapper. Callers from async code should
|
||||
dispatch through asyncio.to_thread() to avoid blocking the event loop."""
|
||||
|
||||
def __init__(self, db_path: str):
|
||||
self._conn = sqlite3.connect(db_path, check_same_thread=False)
|
||||
self._conn.row_factory = sqlite3.Row
|
||||
self._lock = threading.Lock()
|
||||
with self._lock, self._conn:
|
||||
self._conn.execute("PRAGMA journal_mode=WAL")
|
||||
self._conn.executescript(SCHEMA)
|
||||
self._migrate()
|
||||
|
||||
def _migrate(self) -> None:
|
||||
# Older DBs created before message_created_at existed.
|
||||
columns = {row["name"] for row in self._conn.execute("PRAGMA table_info(reactions_log)")}
|
||||
if "message_created_at" not in columns:
|
||||
self._conn.execute("ALTER TABLE reactions_log ADD COLUMN message_created_at TEXT")
|
||||
|
||||
def upsert_target(self, target_id, target_type: str, name: str) -> None:
|
||||
with self._lock, self._conn:
|
||||
self._conn.execute(
|
||||
"""
|
||||
INSERT INTO targets (target_id, target_type, name)
|
||||
VALUES (?, ?, ?)
|
||||
ON CONFLICT(target_id) DO UPDATE SET name = excluded.name
|
||||
""",
|
||||
(str(target_id), target_type, name),
|
||||
)
|
||||
|
||||
def has_reacted(self, message_id, emoji: str) -> bool:
|
||||
with self._lock:
|
||||
cur = self._conn.execute(
|
||||
"SELECT 1 FROM reactions_log WHERE message_id = ? AND emoji = ? LIMIT 1",
|
||||
(str(message_id), emoji),
|
||||
)
|
||||
return cur.fetchone() is not None
|
||||
|
||||
def record_reaction(
|
||||
self, message_id, channel_id, target_id, author_id, emoji: str, source: str, message_created_at=None
|
||||
) -> bool:
|
||||
created_at = message_created_at.strftime("%Y-%m-%dT%H:%M:%fZ") if message_created_at else None
|
||||
with self._lock, self._conn:
|
||||
cur = self._conn.execute(
|
||||
"""
|
||||
INSERT OR IGNORE INTO reactions_log
|
||||
(message_id, channel_id, target_id, author_id, emoji, source, message_created_at)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?)
|
||||
""",
|
||||
(str(message_id), str(channel_id), str(target_id), str(author_id), emoji, source, created_at),
|
||||
)
|
||||
inserted = cur.rowcount > 0
|
||||
if inserted:
|
||||
self._conn.execute(
|
||||
"UPDATE targets SET reaction_count = reaction_count + 1 WHERE target_id = ?",
|
||||
(str(target_id),),
|
||||
)
|
||||
return inserted
|
||||
|
||||
def get_scan_state(self, channel_id) -> Optional[dict]:
|
||||
with self._lock:
|
||||
cur = self._conn.execute(
|
||||
"SELECT * FROM channel_scan_state WHERE channel_id = ?",
|
||||
(str(channel_id),),
|
||||
)
|
||||
row = cur.fetchone()
|
||||
return dict(row) if row else None
|
||||
|
||||
def set_scan_state(self, channel_id, target_id, newest_id_seen, backlog_complete: bool) -> None:
|
||||
with self._lock, self._conn:
|
||||
self._conn.execute(
|
||||
"""
|
||||
INSERT INTO channel_scan_state
|
||||
(channel_id, target_id, newest_id_seen, backlog_complete, last_scanned_at)
|
||||
VALUES (?, ?, ?, ?, strftime('%Y-%m-%dT%H:%M:%fZ','now'))
|
||||
ON CONFLICT(channel_id) DO UPDATE SET
|
||||
newest_id_seen = excluded.newest_id_seen,
|
||||
backlog_complete = excluded.backlog_complete,
|
||||
last_scanned_at = excluded.last_scanned_at
|
||||
""",
|
||||
(str(channel_id), str(target_id), str(newest_id_seen) if newest_id_seen else None, int(backlog_complete)),
|
||||
)
|
||||
|
||||
def get_all_reactions(self) -> list:
|
||||
with self._lock:
|
||||
cur = self._conn.execute(
|
||||
"SELECT id, message_id, channel_id, target_id, emoji FROM reactions_log ORDER BY id"
|
||||
)
|
||||
return [dict(row) for row in cur.fetchall()]
|
||||
|
||||
def delete_reaction(self, row_id: int) -> None:
|
||||
with self._lock, self._conn:
|
||||
row = self._conn.execute(
|
||||
"SELECT target_id FROM reactions_log WHERE id = ?", (row_id,)
|
||||
).fetchone()
|
||||
self._conn.execute("DELETE FROM reactions_log WHERE id = ?", (row_id,))
|
||||
if row:
|
||||
self._conn.execute(
|
||||
"UPDATE targets SET reaction_count = MAX(reaction_count - 1, 0) WHERE target_id = ?",
|
||||
(row["target_id"],),
|
||||
)
|
||||
|
||||
def get_stats(self) -> list:
|
||||
with self._lock:
|
||||
cur = self._conn.execute(
|
||||
"""
|
||||
SELECT t.target_id, t.target_type, t.name, t.reaction_count,
|
||||
COUNT(r.id) AS logged_reactions, MAX(r.message_created_at) AS last_reaction_at
|
||||
FROM targets t
|
||||
LEFT JOIN reactions_log r ON r.target_id = t.target_id
|
||||
GROUP BY t.target_id
|
||||
ORDER BY t.reaction_count DESC
|
||||
"""
|
||||
)
|
||||
return [dict(row) for row in cur.fetchall()]
|
||||
|
||||
def close(self) -> None:
|
||||
with self._lock:
|
||||
self._conn.close()
|
||||
46
bot/undo.py
Normal file
46
bot/undo.py
Normal file
@@ -0,0 +1,46 @@
|
||||
import asyncio
|
||||
import logging
|
||||
from typing import Optional
|
||||
|
||||
import discord
|
||||
|
||||
from bot.display import StatsDisplay
|
||||
from bot.reactor import human_jitter
|
||||
from bot.storage import Storage
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
|
||||
async def undo_all_reactions(client: discord.Client, storage: Storage, display: Optional[StatsDisplay] = None) -> None:
|
||||
"""Remove every reaction this script has previously logged, then clear those log entries."""
|
||||
rows = await asyncio.to_thread(storage.get_all_reactions)
|
||||
total = len(rows)
|
||||
log.info("Undo: removing %d previously-added reaction(s)...", total)
|
||||
|
||||
removed = 0
|
||||
for i, row in enumerate(rows, start=1):
|
||||
channel_id = int(row["channel_id"])
|
||||
message_id = int(row["message_id"])
|
||||
emoji = row["emoji"]
|
||||
should_clear_log = True
|
||||
|
||||
try:
|
||||
channel = client.get_channel(channel_id) or await client.fetch_channel(channel_id)
|
||||
message = await channel.fetch_message(message_id)
|
||||
await message.remove_reaction(emoji, client.user)
|
||||
removed += 1
|
||||
log.info("[%d/%d] Removed reaction from message %s", i, total, message_id)
|
||||
except discord.NotFound:
|
||||
log.info("[%d/%d] Message/reaction %s already gone, clearing log entry", i, total, message_id)
|
||||
except discord.HTTPException:
|
||||
log.exception("[%d/%d] Failed to remove reaction from message %s, will retry next run", i, total, message_id)
|
||||
should_clear_log = False
|
||||
|
||||
if should_clear_log:
|
||||
await asyncio.to_thread(storage.delete_reaction, row["id"])
|
||||
if display is not None:
|
||||
await display.refresh()
|
||||
|
||||
await asyncio.sleep(human_jitter())
|
||||
|
||||
log.info("Undo complete: removed %d/%d reaction(s)", removed, total)
|
||||
Reference in New Issue
Block a user