Compare commits
32 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| ff78f7d5bf | |||
| 384820370a | |||
| 52cbd68778 | |||
| 447672d972 | |||
| a098106468 | |||
| 8ccd92633c | |||
| f103dcf4d4 | |||
| b931f05fa1 | |||
| 357ea51066 | |||
| e4534bcace | |||
| e2ae5bd885 | |||
| 1eb9a32683 | |||
| dba0930e3e | |||
| 4350236ea7 | |||
| 8339921f93 | |||
| f432aa2b9d | |||
| 85234404c3 | |||
| 407938bf43 | |||
| 5a435b6a5d | |||
| b74c7cda48 | |||
| daa9eb671b | |||
| cf363fd738 | |||
| 9a0d4090f5 | |||
| d05dc81667 | |||
| a4dd540f44 | |||
| 63e5dd1796 | |||
| daa370e09f | |||
| 58c0873987 | |||
| 5de3fd77bf | |||
| 873def8456 | |||
| c3c8baa4b2 | |||
| 850c5d7abb |
+21
@@ -0,0 +1,21 @@
|
|||||||
|
[submodule "src/modules/voicefix"]
|
||||||
|
path = src/modules/voicefix
|
||||||
|
url = https://github.com/Intery/StudyLion-voicefix.git
|
||||||
|
[submodule "src/modules/streamalerts"]
|
||||||
|
path = src/modules/streamalerts
|
||||||
|
url = https://github.com/Intery/StudyLion-streamalerts.git
|
||||||
|
[submodule "src/modules/messagelogger"]
|
||||||
|
path = src/modules/messagelogger
|
||||||
|
url = https://git.thewisewolf.dev/HoloTech/discord-messagelogger-plugin.git
|
||||||
|
[submodule "src/modules/voicelog"]
|
||||||
|
path = src/modules/voicelog
|
||||||
|
url = https://git.thewisewolf.dev/HoloTech/voicelog-plugin.git
|
||||||
|
[submodule "src/data"]
|
||||||
|
path = src/data
|
||||||
|
url = https://git.thewisewolf.dev/HoloTech/psqlmapper.git
|
||||||
|
[submodule "src/modules/profiles"]
|
||||||
|
path = src/modules/profiles
|
||||||
|
url = https://git.thewisewolf.dev/HoloTech/profiles-plugin.git
|
||||||
|
[submodule "src/modules/pluscampaign"]
|
||||||
|
path = src/modules/pluscampaign
|
||||||
|
url = https://git.thewisewolf.dev/CarmiCoven/pluscampaign-plugin.git
|
||||||
+29
-25
@@ -1,10 +1,14 @@
|
|||||||
|
BEGIN;
|
||||||
|
|
||||||
-- Metadata {{{
|
-- Metadata {{{
|
||||||
CREATE TABLE VersionHistory(
|
CREATE TABLE version_history(
|
||||||
version INTEGER NOT NULL,
|
component TEXT NOT NULL,
|
||||||
time TIMESTAMP WITH TIME ZONE DEFAULT CURRENT_TIMESTAMP NOT NULL,
|
from_version INTEGER NOT NULL,
|
||||||
author TEXT
|
to_version INTEGER NOT NULL,
|
||||||
|
author TEXT NOT NULL,
|
||||||
|
_timestamp TIMESTAMPTZ NOT NULL DEFAULT NOW()
|
||||||
);
|
);
|
||||||
INSERT INTO VersionHistory (version, author) VALUES (1, 'Initial Creation');
|
INSERT INTO version_history (component, from_version, to_version, author) VALUES ('ROOT', 0, 1, 'Initial Creation');
|
||||||
|
|
||||||
|
|
||||||
CREATE OR REPLACE FUNCTION update_timestamp_column()
|
CREATE OR REPLACE FUNCTION update_timestamp_column()
|
||||||
@@ -14,6 +18,24 @@ BEGIN
|
|||||||
RETURN NEW;
|
RETURN NEW;
|
||||||
END;
|
END;
|
||||||
$$ language 'plpgsql';
|
$$ language 'plpgsql';
|
||||||
|
|
||||||
|
|
||||||
|
CREATE OR REPLACE FUNCTION current_module_version(module_name TEXT)
|
||||||
|
RETURNS INTEGER
|
||||||
|
AS $$
|
||||||
|
SELECT
|
||||||
|
to_version
|
||||||
|
FROM version_history
|
||||||
|
WHERE
|
||||||
|
component = $1
|
||||||
|
ORDER BY _timestamp DESC
|
||||||
|
LIMIT 1;
|
||||||
|
$$ LANGUAGE SQL;
|
||||||
|
|
||||||
|
CREATE TABLE app_config(
|
||||||
|
appname TEXT PRIMARY KEY,
|
||||||
|
created_at TIMESTAMPTZ NOT NULL DEFAULT now()
|
||||||
|
);
|
||||||
-- }}}
|
-- }}}
|
||||||
|
|
||||||
-- App metadata {{{
|
-- App metadata {{{
|
||||||
@@ -31,26 +53,8 @@ CREATE TABLE bot_config(
|
|||||||
);
|
);
|
||||||
-- }}}
|
-- }}}
|
||||||
|
|
||||||
-- Channel Linker {{{
|
-- TODO: Profile data
|
||||||
|
|
||||||
CREATE TABLE links(
|
|
||||||
linkid SERIAL PRIMARY KEY,
|
|
||||||
name TEXT
|
|
||||||
);
|
|
||||||
|
|
||||||
CREATE TABLE channel_webhooks(
|
|
||||||
channelid BIGINT PRIMARY KEY,
|
|
||||||
webhookid BIGINT NOT NULL,
|
|
||||||
token TEXT NOT NULL
|
|
||||||
);
|
|
||||||
|
|
||||||
CREATE TABLE channel_links(
|
|
||||||
linkid INTEGER NOT NULL REFERENCES links (linkid) ON DELETE CASCADE,
|
|
||||||
channelid BIGINT NOT NULL REFERENCES channel_webhooks (channelid) ON DELETE CASCADE,
|
|
||||||
PRIMARY KEY (linkid, channelid)
|
|
||||||
);
|
|
||||||
|
|
||||||
|
|
||||||
-- }}}
|
COMMIT;
|
||||||
|
|
||||||
-- vim: set fdm=marker:
|
-- vim: set fdm=marker:
|
||||||
|
|||||||
+6
-5
@@ -1,7 +1,8 @@
|
|||||||
aiohttp==3.7.4.post0
|
aiohttp
|
||||||
cachetools==4.2.2
|
cachetools
|
||||||
configparser==5.0.2
|
configparser
|
||||||
discord.py [voice]
|
discord.py [voice]
|
||||||
iso8601==0.1.16
|
iso8601
|
||||||
psycopg[pool]
|
psycopg[pool]
|
||||||
pytz==2021.1
|
pytz
|
||||||
|
twitchAPI
|
||||||
|
|||||||
@@ -0,0 +1,3 @@
|
|||||||
|
from .translator import SOURCE_LOCALE, LeoBabel, LocalBabel, LazyStr, ctx_locale, ctx_translator
|
||||||
|
|
||||||
|
babel = LocalBabel('babel')
|
||||||
@@ -0,0 +1,81 @@
|
|||||||
|
from enum import Enum
|
||||||
|
from . import babel
|
||||||
|
|
||||||
|
_p = babel._p
|
||||||
|
|
||||||
|
|
||||||
|
class LocaleMap(Enum):
|
||||||
|
american_english = 'en-US'
|
||||||
|
british_english = 'en-GB'
|
||||||
|
bulgarian = 'bg'
|
||||||
|
chinese = 'zh-CN'
|
||||||
|
taiwan_chinese = 'zh-TW'
|
||||||
|
croatian = 'hr'
|
||||||
|
czech = 'cs'
|
||||||
|
danish = 'da'
|
||||||
|
dutch = 'nl'
|
||||||
|
finnish = 'fi'
|
||||||
|
french = 'fr'
|
||||||
|
german = 'de'
|
||||||
|
greek = 'el'
|
||||||
|
hindi = 'hi'
|
||||||
|
hungarian = 'hu'
|
||||||
|
italian = 'it'
|
||||||
|
japanese = 'ja'
|
||||||
|
korean = 'ko'
|
||||||
|
lithuanian = 'lt'
|
||||||
|
norwegian = 'no'
|
||||||
|
polish = 'pl'
|
||||||
|
brazil_portuguese = 'pt-BR'
|
||||||
|
romanian = 'ro'
|
||||||
|
russian = 'ru'
|
||||||
|
spain_spanish = 'es-ES'
|
||||||
|
swedish = 'sv-SE'
|
||||||
|
thai = 'th'
|
||||||
|
turkish = 'tr'
|
||||||
|
ukrainian = 'uk'
|
||||||
|
vietnamese = 'vi'
|
||||||
|
hebrew = 'he-IL'
|
||||||
|
|
||||||
|
|
||||||
|
# Original Discord names
|
||||||
|
locale_names = {
|
||||||
|
'id': (_p('localenames|locale:id', "Indonesian"), "Bahasa Indonesia"),
|
||||||
|
'da': (_p('localenames|locale:da', "Danish"), "Dansk"),
|
||||||
|
'de': (_p('localenames|locale:de', "German"), "Deutsch"),
|
||||||
|
'en-GB': (_p('localenames|locale:en-GB', "English, UK"), "English, UK"),
|
||||||
|
'en-US': (_p('localenames|locale:en-US', "English, US"), "English, US"),
|
||||||
|
'es-ES': (_p('localenames|locale:es-ES', "Spanish"), "Español"),
|
||||||
|
'fr': (_p('localenames|locale:fr', "French"), "Français"),
|
||||||
|
'hr': (_p('localenames|locale:hr', "Croatian"), "Hrvatski"),
|
||||||
|
'it': (_p('localenames|locale:it', "Italian"), "Italiano"),
|
||||||
|
'lt': (_p('localenames|locale:lt', "Lithuanian"), "Lietuviškai"),
|
||||||
|
'hu': (_p('localenames|locale:hu', "Hungarian"), "Magyar"),
|
||||||
|
'nl': (_p('localenames|locale:nl', "Dutch"), "Nederlands"),
|
||||||
|
'no': (_p('localenames|locale:no', "Norwegian"), "Norsk"),
|
||||||
|
'pl': (_p('localenames|locale:pl', "Polish"), "Polski"),
|
||||||
|
'pt-BR': (_p('localenames|locale:pt-BR', "Portuguese, Brazilian"), "Português do Brasil"),
|
||||||
|
'ro': (_p('localenames|locale:ro', "Romanian, Romania"), "Română"),
|
||||||
|
'fi': (_p('localenames|locale:fi', "Finnish"), "Suomi"),
|
||||||
|
'sv-SE': (_p('localenames|locale:sv-SE', "Swedish"), "Svenska"),
|
||||||
|
'vi': (_p('localenames|locale:vi', "Vietnamese"), "Tiếng Việt"),
|
||||||
|
'tr': (_p('localenames|locale:tr', "Turkish"), "Türkçe"),
|
||||||
|
'cs': (_p('localenames|locale:cs', "Czech"), "Čeština"),
|
||||||
|
'el': (_p('localenames|locale:el', "Greek"), "Ελληνικά"),
|
||||||
|
'bg': (_p('localenames|locale:bg', "Bulgarian"), "български"),
|
||||||
|
'ru': (_p('localenames|locale:ru', "Russian"), "Pусский"),
|
||||||
|
'uk': (_p('localenames|locale:uk', "Ukrainian"), "Українська"),
|
||||||
|
'hi': (_p('localenames|locale:hi', "Hindi"), "हिन्दी"),
|
||||||
|
'th': (_p('localenames|locale:th', "Thai"), "ไทย"),
|
||||||
|
'zh-CN': (_p('localenames|locale:zh-CN', "Chinese, China"), "中文"),
|
||||||
|
'ja': (_p('localenames|locale:ja', "Japanese"), "日本語"),
|
||||||
|
'zh-TW': (_p('localenames|locale:zh-TW', "Chinese, Taiwan"), "繁體中文"),
|
||||||
|
'ko': (_p('localenames|locale:ko', "Korean"), "한국어"),
|
||||||
|
}
|
||||||
|
|
||||||
|
# More names for languages not supported by Discord
|
||||||
|
locale_names |= {
|
||||||
|
'he': (_p('localenames|locale:he', "Hebrew"), "Hebrew"),
|
||||||
|
'he-IL': (_p('localenames|locale:he-IL', "Hebrew"), "Hebrew"),
|
||||||
|
'ceaser': (_p('localenames|locale:test', "Test Language"), "dfbtfs"),
|
||||||
|
}
|
||||||
@@ -0,0 +1,108 @@
|
|||||||
|
from typing import Optional
|
||||||
|
import logging
|
||||||
|
from contextvars import ContextVar
|
||||||
|
from collections import defaultdict
|
||||||
|
from enum import Enum
|
||||||
|
|
||||||
|
import gettext
|
||||||
|
|
||||||
|
from discord.app_commands import Translator, locale_str
|
||||||
|
from discord.enums import Locale
|
||||||
|
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
SOURCE_LOCALE = 'en_GB'
|
||||||
|
ctx_locale: ContextVar[str] = ContextVar('locale', default=SOURCE_LOCALE)
|
||||||
|
ctx_translator: ContextVar['LeoBabel'] = ContextVar('translator', default=None) # type: ignore
|
||||||
|
|
||||||
|
null = gettext.NullTranslations()
|
||||||
|
|
||||||
|
|
||||||
|
class LeoBabel(Translator):
|
||||||
|
def __init__(self):
|
||||||
|
self.supported_locales = {loc.name for loc in Locale}
|
||||||
|
self.supported_domains = {}
|
||||||
|
self.translators = defaultdict(dict) # locale -> domain -> GNUTranslator
|
||||||
|
|
||||||
|
async def load(self):
|
||||||
|
pass
|
||||||
|
|
||||||
|
async def unload(self):
|
||||||
|
self.translators.clear()
|
||||||
|
|
||||||
|
def get_translator(self, locale: Optional[str], domain):
|
||||||
|
return null
|
||||||
|
|
||||||
|
def t(self, lazystr, locale=None):
|
||||||
|
return lazystr._translate_with(null)
|
||||||
|
|
||||||
|
async def translate(self, string: locale_str, locale: Locale, context):
|
||||||
|
if not isinstance(string, LazyStr):
|
||||||
|
return string
|
||||||
|
else:
|
||||||
|
return string.message
|
||||||
|
|
||||||
|
ctx_translator.set(LeoBabel())
|
||||||
|
|
||||||
|
class Method(Enum):
|
||||||
|
GETTEXT = 'gettext'
|
||||||
|
NGETTEXT = 'ngettext'
|
||||||
|
PGETTEXT = 'pgettext'
|
||||||
|
NPGETTEXT = 'npgettext'
|
||||||
|
|
||||||
|
|
||||||
|
class LocalBabel:
|
||||||
|
def __init__(self, domain):
|
||||||
|
self.domain = domain
|
||||||
|
|
||||||
|
@property
|
||||||
|
def methods(self):
|
||||||
|
return (self._, self._n, self._p, self._np)
|
||||||
|
|
||||||
|
def _(self, message):
|
||||||
|
return LazyStr(Method.GETTEXT, message, domain=self.domain)
|
||||||
|
|
||||||
|
def _n(self, singular, plural, n):
|
||||||
|
return LazyStr(Method.NGETTEXT, singular, plural, n, domain=self.domain)
|
||||||
|
|
||||||
|
def _p(self, context, message):
|
||||||
|
return LazyStr(Method.PGETTEXT, context, message, domain=self.domain)
|
||||||
|
|
||||||
|
def _np(self, context, singular, plural, n):
|
||||||
|
return LazyStr(Method.NPGETTEXT, context, singular, plural, n, domain=self.domain)
|
||||||
|
|
||||||
|
|
||||||
|
class LazyStr(locale_str):
|
||||||
|
__slots__ = ('method', 'args', 'domain', 'locale')
|
||||||
|
|
||||||
|
def __init__(self, method, *args, locale=None, domain=None):
|
||||||
|
self.method = method
|
||||||
|
self.args = args
|
||||||
|
self.domain = domain
|
||||||
|
self.locale = locale
|
||||||
|
|
||||||
|
@property
|
||||||
|
def message(self):
|
||||||
|
return self._translate_with(null)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def extras(self):
|
||||||
|
return {'locale': self.locale, 'domain': self.domain}
|
||||||
|
|
||||||
|
def __str__(self):
|
||||||
|
return self.message
|
||||||
|
|
||||||
|
def _translate_with(self, translator: gettext.GNUTranslations):
|
||||||
|
method = getattr(translator, self.method.value)
|
||||||
|
return method(*self.args)
|
||||||
|
|
||||||
|
def __repr__(self) -> str:
|
||||||
|
return f'{self.__class__.__name__}({self.method}, {self.args!r}, locale={self.locale}, domain={self.domain})'
|
||||||
|
|
||||||
|
def __eq__(self, obj: object) -> bool:
|
||||||
|
return isinstance(obj, locale_str) and self.message == obj.message
|
||||||
|
|
||||||
|
def __hash__(self) -> int:
|
||||||
|
return hash(self.args)
|
||||||
@@ -0,0 +1,20 @@
|
|||||||
|
from .translator import ctx_translator
|
||||||
|
from . import babel
|
||||||
|
|
||||||
|
_, _p, _np = babel._, babel._p, babel._np
|
||||||
|
|
||||||
|
|
||||||
|
MONTHS = _p(
|
||||||
|
'utils|months',
|
||||||
|
"January,February,March,April,May,June,July,August,September,October,November,December"
|
||||||
|
)
|
||||||
|
|
||||||
|
SHORT_MONTHS = _p(
|
||||||
|
'utils|short_months',
|
||||||
|
"Jan,Feb,Mar,Apr,May,Jun,Jul,Aug,Sep,Oct,Nov,Dec"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def local_month(month, short=False):
|
||||||
|
string = MONTHS if not short else SHORT_MONTHS
|
||||||
|
return ctx_translator.get().t(string).split(',')[month-1]
|
||||||
+3
-10
@@ -13,8 +13,6 @@ from meta.monitor import ComponentMonitor, StatusLevel, ComponentStatus
|
|||||||
|
|
||||||
from data import Database
|
from data import Database
|
||||||
|
|
||||||
from constants import DATA_VERSION
|
|
||||||
|
|
||||||
|
|
||||||
for name in conf.config.options('LOGGING_LEVELS', no_defaults=True):
|
for name in conf.config.options('LOGGING_LEVELS', no_defaults=True):
|
||||||
logging.getLogger(name).setLevel(conf.logging_levels[name])
|
logging.getLogger(name).setLevel(conf.logging_levels[name])
|
||||||
@@ -54,18 +52,13 @@ async def main():
|
|||||||
intents = discord.Intents.all()
|
intents = discord.Intents.all()
|
||||||
intents.members = True
|
intents.members = True
|
||||||
intents.message_content = True
|
intents.message_content = True
|
||||||
intents.presences = False
|
intents.presences = True
|
||||||
|
|
||||||
async with db.open():
|
async with db.open():
|
||||||
version = await db.version()
|
|
||||||
if version.version != DATA_VERSION:
|
|
||||||
error = f"Data model version is {version}, required version is {DATA_VERSION}! Please migrate."
|
|
||||||
logger.critical(error)
|
|
||||||
raise RuntimeError(error)
|
|
||||||
|
|
||||||
async with aiohttp.ClientSession() as session:
|
async with aiohttp.ClientSession() as session:
|
||||||
async with LionBot(
|
async with LionBot(
|
||||||
command_prefix='!leo!',
|
command_prefix=conf.bot.get('prefix', '!!'),
|
||||||
intents=intents,
|
intents=intents,
|
||||||
appname=appname,
|
appname=appname,
|
||||||
shardname=shardname,
|
shardname=shardname,
|
||||||
@@ -81,7 +74,7 @@ async def main():
|
|||||||
shard_count=sharding.shard_count,
|
shard_count=sharding.shard_count,
|
||||||
help_command=None,
|
help_command=None,
|
||||||
proxy=conf.bot.get('proxy', None),
|
proxy=conf.bot.get('proxy', None),
|
||||||
chunk_guilds_at_startup=False,
|
chunk_guilds_at_startup=True,
|
||||||
) as lionbot:
|
) as lionbot:
|
||||||
ctx_bot.set(lionbot)
|
ctx_bot.set(lionbot)
|
||||||
lionbot.system_monitor.add_component(
|
lionbot.system_monitor.add_component(
|
||||||
|
|||||||
@@ -0,0 +1,26 @@
|
|||||||
|
from data import Registry, RowModel, Table
|
||||||
|
from data.columns import String, Timestamp, Integer, Bool
|
||||||
|
|
||||||
|
|
||||||
|
class VersionHistory(RowModel):
|
||||||
|
"""
|
||||||
|
CREATE TABLE version_history(
|
||||||
|
component TEXT NOT NULL,
|
||||||
|
from_version INTEGER NOT NULL,
|
||||||
|
to_version INTEGER NOT NULL,
|
||||||
|
author TEXT NOT NULL,
|
||||||
|
_timestamp TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
||||||
|
);
|
||||||
|
"""
|
||||||
|
_tablename_ = 'version_history'
|
||||||
|
_cache_ = {}
|
||||||
|
|
||||||
|
component = String()
|
||||||
|
from_version = Integer()
|
||||||
|
to_version = Integer()
|
||||||
|
author = String()
|
||||||
|
_timestamp = Timestamp()
|
||||||
|
|
||||||
|
|
||||||
|
class BotData(Registry):
|
||||||
|
version_history = VersionHistory.table
|
||||||
+4
-3
@@ -1,6 +1,7 @@
|
|||||||
CONFIG_FILE = "config/bot.conf"
|
CONFIG_FILE = "config/bot.conf"
|
||||||
DATA_VERSION = 1
|
|
||||||
|
|
||||||
MAX_COINS = 2147483647 - 1
|
|
||||||
|
|
||||||
HINT_ICON = "https://projects.iamcal.com/emoji-data/img-apple-64/1f4a1.png"
|
HINT_ICON = "https://projects.iamcal.com/emoji-data/img-apple-64/1f4a1.png"
|
||||||
|
|
||||||
|
SCHEMA_VERSIONS = {
|
||||||
|
'ROOT': 1,
|
||||||
|
}
|
||||||
|
|||||||
@@ -1,4 +1,6 @@
|
|||||||
|
from babel import LocalBabel
|
||||||
|
|
||||||
|
babel = LocalBabel('core')
|
||||||
|
|
||||||
async def setup(bot):
|
async def setup(bot):
|
||||||
from .cog import CoreCog
|
from .cog import CoreCog
|
||||||
|
|||||||
@@ -0,0 +1,227 @@
|
|||||||
|
"""
|
||||||
|
Additional abstract setting types useful for StudyLion settings.
|
||||||
|
"""
|
||||||
|
from typing import Optional
|
||||||
|
import json
|
||||||
|
import traceback
|
||||||
|
|
||||||
|
import discord
|
||||||
|
from discord.enums import TextStyle
|
||||||
|
|
||||||
|
from settings.base import ParentID
|
||||||
|
from settings.setting_types import IntegerSetting, StringSetting
|
||||||
|
from meta import conf
|
||||||
|
from meta.errors import UserInputError
|
||||||
|
from babel.translator import ctx_translator
|
||||||
|
from utils.lib import MessageArgs
|
||||||
|
|
||||||
|
from . import babel
|
||||||
|
|
||||||
|
_p = babel._p
|
||||||
|
|
||||||
|
|
||||||
|
class MessageSetting(StringSetting):
|
||||||
|
"""
|
||||||
|
Typed Setting ABC representing a message sent to Discord.
|
||||||
|
|
||||||
|
Data is a json-formatted string dict with at least one of the fields 'content', 'embed', 'embeds'
|
||||||
|
Value is the corresponding dictionary
|
||||||
|
"""
|
||||||
|
# TODO: Extend to support format keys
|
||||||
|
|
||||||
|
_accepts = _p(
|
||||||
|
'settype:message|accepts',
|
||||||
|
"JSON formatted raw message data"
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
async def download_attachment(attached: discord.Attachment):
|
||||||
|
"""
|
||||||
|
Download a discord.Attachment with some basic filetype and file size validation.
|
||||||
|
"""
|
||||||
|
t = ctx_translator.get().t
|
||||||
|
|
||||||
|
error = None
|
||||||
|
decoded = None
|
||||||
|
if attached.content_type and not ('json' in attached.content_type):
|
||||||
|
error = t(_p(
|
||||||
|
'settype:message|download|error:not_json',
|
||||||
|
"The attached message data is not a JSON file!"
|
||||||
|
))
|
||||||
|
elif attached.size > 10000:
|
||||||
|
error = t(_p(
|
||||||
|
'settype:message|download|error:size',
|
||||||
|
"The attached message data is too large!"
|
||||||
|
))
|
||||||
|
else:
|
||||||
|
content = await attached.read()
|
||||||
|
try:
|
||||||
|
decoded = content.decode('UTF-8')
|
||||||
|
except UnicodeDecodeError:
|
||||||
|
error = t(_p(
|
||||||
|
'settype:message|download|error:decoding',
|
||||||
|
"Could not decode the message data. Please ensure it is saved with the `UTF-8` encoding."
|
||||||
|
))
|
||||||
|
|
||||||
|
if error is not None:
|
||||||
|
raise UserInputError(error)
|
||||||
|
else:
|
||||||
|
return decoded
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def value_to_args(cls, parent_id: ParentID, value: dict, **kwargs) -> MessageArgs:
|
||||||
|
if not value:
|
||||||
|
return None
|
||||||
|
|
||||||
|
args = {}
|
||||||
|
args['content'] = value.get('content', "")
|
||||||
|
if 'embed' in value:
|
||||||
|
embed = discord.Embed.from_dict(value['embed'])
|
||||||
|
args['embed'] = embed
|
||||||
|
if 'embeds' in value:
|
||||||
|
embeds = []
|
||||||
|
for embed_data in value['embeds']:
|
||||||
|
embeds.append(discord.Embed.from_dict(embed_data))
|
||||||
|
args['embeds'] = embeds
|
||||||
|
return MessageArgs(**args)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _data_from_value(cls, parent_id: ParentID, value: Optional[dict], **kwargs):
|
||||||
|
if value and any(value.get(key, None) for key in ('content', 'embed', 'embeds')):
|
||||||
|
data = json.dumps(value)
|
||||||
|
else:
|
||||||
|
data = None
|
||||||
|
return data
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _data_to_value(cls, parent_id: ParentID, data: Optional[str], **kwargs):
|
||||||
|
if data:
|
||||||
|
value = json.loads(data)
|
||||||
|
else:
|
||||||
|
value = None
|
||||||
|
return value
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
async def _parse_string(cls, parent_id: ParentID, string: str, **kwargs):
|
||||||
|
"""
|
||||||
|
Provided user string can be downright random.
|
||||||
|
|
||||||
|
If it isn't json-formatted, treat it as the content of the message.
|
||||||
|
If it is, do basic checking on the length and embeds.
|
||||||
|
"""
|
||||||
|
string = string.strip()
|
||||||
|
if not string or string.lower() == 'none':
|
||||||
|
return None
|
||||||
|
|
||||||
|
t = ctx_translator.get().t
|
||||||
|
|
||||||
|
error_tip = t(_p(
|
||||||
|
'settype:message|error_suffix',
|
||||||
|
"You can view, test, and fix your embed using the online [embed builder]({link})."
|
||||||
|
)).format(
|
||||||
|
link="https://glitchii.github.io/embedbuilder/?editor=json"
|
||||||
|
)
|
||||||
|
|
||||||
|
if string.startswith('{') and string.endswith('}'):
|
||||||
|
# Assume the string is a json-formatted message dict
|
||||||
|
try:
|
||||||
|
value = json.loads(string)
|
||||||
|
except json.JSONDecodeError as err:
|
||||||
|
error = t(_p(
|
||||||
|
'settype:message|error:invalid_json',
|
||||||
|
"The provided message data was not a valid JSON document!\n"
|
||||||
|
"`{error}`"
|
||||||
|
)).format(error=str(err))
|
||||||
|
raise UserInputError(error + '\n' + error_tip)
|
||||||
|
|
||||||
|
if not isinstance(value, dict) or not any(value.get(key, None) for key in ('content', 'embed', 'embeds')):
|
||||||
|
error = t(_p(
|
||||||
|
'settype:message|error:json_missing_keys',
|
||||||
|
"Message data must be a JSON object with at least one of the following fields: "
|
||||||
|
"`content`, `embed`, `embeds`"
|
||||||
|
))
|
||||||
|
raise UserInputError(error + '\n' + error_tip)
|
||||||
|
|
||||||
|
embed_data = value.get('embed', None)
|
||||||
|
if not isinstance(embed_data, dict):
|
||||||
|
error = t(_p(
|
||||||
|
'settype:message|error:json_embed_type',
|
||||||
|
"`embed` field must be a valid JSON object."
|
||||||
|
))
|
||||||
|
raise UserInputError(error + '\n' + error_tip)
|
||||||
|
|
||||||
|
embeds_data = value.get('embeds', [])
|
||||||
|
if not isinstance(embeds_data, list):
|
||||||
|
error = t(_p(
|
||||||
|
'settype:message|error:json_embeds_type',
|
||||||
|
"`embeds` field must be a list."
|
||||||
|
))
|
||||||
|
raise UserInputError(error + '\n' + error_tip)
|
||||||
|
|
||||||
|
if embed_data and embeds_data:
|
||||||
|
error = t(_p(
|
||||||
|
'settype:message|error:json_embed_embeds',
|
||||||
|
"Message data cannot include both `embed` and `embeds`."
|
||||||
|
))
|
||||||
|
raise UserInputError(error + '\n' + error_tip)
|
||||||
|
|
||||||
|
content_data = value.get('content', "")
|
||||||
|
if not isinstance(content_data, str):
|
||||||
|
error = t(_p(
|
||||||
|
'settype:message|error:json_content_type',
|
||||||
|
"`content` field must be a string."
|
||||||
|
))
|
||||||
|
raise UserInputError(error + '\n' + error_tip)
|
||||||
|
|
||||||
|
# Validate embeds, which is the most likely place for something to go wrong
|
||||||
|
embeds = [embed_data] if embed_data else embeds_data
|
||||||
|
try:
|
||||||
|
for embed in embeds:
|
||||||
|
discord.Embed.from_dict(embed)
|
||||||
|
except Exception as e:
|
||||||
|
# from_dict may raise a range of possible exceptions.
|
||||||
|
raw_error = ''.join(
|
||||||
|
traceback.TracebackException.from_exception(e).format_exception_only()
|
||||||
|
)
|
||||||
|
error = t(_p(
|
||||||
|
'ui:settype:message|error:embed_conversion',
|
||||||
|
"Could not parse the message embed data.\n"
|
||||||
|
"**Error:** `{exception}`"
|
||||||
|
)).format(exception=raw_error)
|
||||||
|
raise UserInputError(error + '\n' + error_tip)
|
||||||
|
|
||||||
|
# At this point, the message will at least successfully convert into MessageArgs
|
||||||
|
# There are numerous ways it could still be invalid, e.g. invalid urls, or too-long fields
|
||||||
|
# or the total message content being too long, or too many fields, etc
|
||||||
|
# This will need to be caught in anything which displays a message parsed from user data.
|
||||||
|
else:
|
||||||
|
# Either the string is not json formatted, or the formatting is broken
|
||||||
|
# Assume the string is a content message
|
||||||
|
value = {
|
||||||
|
'content': string
|
||||||
|
}
|
||||||
|
return json.dumps(value)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _format_data(cls, parent_id: ParentID, data: Optional[str], **kwargs):
|
||||||
|
if not data:
|
||||||
|
return None
|
||||||
|
|
||||||
|
value = cls._data_to_value(parent_id, data, **kwargs)
|
||||||
|
content = value.get('content', "")
|
||||||
|
if 'embed' in value or 'embeds' in value or len(content) > 100:
|
||||||
|
t = ctx_translator.get().t
|
||||||
|
formatted = t(_p(
|
||||||
|
'settype:message|format:too_long',
|
||||||
|
"Too long to display! See Preview."
|
||||||
|
))
|
||||||
|
else:
|
||||||
|
formatted = content
|
||||||
|
|
||||||
|
return formatted
|
||||||
|
|
||||||
|
@property
|
||||||
|
def input_field(self):
|
||||||
|
field = super().input_field
|
||||||
|
field.style = TextStyle.long
|
||||||
|
return field
|
||||||
Submodule
+1
Submodule src/data added at c495b4d097
@@ -1,9 +0,0 @@
|
|||||||
from .conditions import Condition, condition, NULL
|
|
||||||
from .database import Database
|
|
||||||
from .models import RowModel, RowTable, WeakCache
|
|
||||||
from .table import Table
|
|
||||||
from .base import Expression, RawExpr
|
|
||||||
from .columns import ColumnExpr, Column, Integer, String
|
|
||||||
from .registry import Registry, AttachableClass, Attachable
|
|
||||||
from .adapted import RegisterEnum
|
|
||||||
from .queries import ORDER, NULLS, JOINTYPE
|
|
||||||
@@ -1,40 +0,0 @@
|
|||||||
# from enum import Enum
|
|
||||||
from typing import Optional
|
|
||||||
from psycopg.types.enum import register_enum, EnumInfo
|
|
||||||
from psycopg import AsyncConnection
|
|
||||||
from .registry import Attachable, Registry
|
|
||||||
|
|
||||||
|
|
||||||
class RegisterEnum(Attachable):
|
|
||||||
def __init__(self, enum, name: Optional[str] = None, mapper=None):
|
|
||||||
super().__init__()
|
|
||||||
self.enum = enum
|
|
||||||
self.name = name or enum.__name__
|
|
||||||
self.mapping = mapper(enum) if mapper is not None else self._mapper()
|
|
||||||
|
|
||||||
def _mapper(self):
|
|
||||||
return {m: m.value[0] for m in self.enum}
|
|
||||||
|
|
||||||
def attach_to(self, registry: Registry):
|
|
||||||
self._registry = registry
|
|
||||||
registry.init_task(self.on_init)
|
|
||||||
return self
|
|
||||||
|
|
||||||
async def on_init(self, registry: Registry):
|
|
||||||
connector = registry._conn
|
|
||||||
if connector is None:
|
|
||||||
raise ValueError("Cannot initialise without connector!")
|
|
||||||
connector.connect_hook(self.connection_hook)
|
|
||||||
# await connector.refresh_pool()
|
|
||||||
# The below may be somewhat dangerous
|
|
||||||
# But adaption should never write to the database
|
|
||||||
await connector.map_over_pool(self.connection_hook)
|
|
||||||
# if conn := connector.conn:
|
|
||||||
# # Ensure the adaption is run in the current context as well
|
|
||||||
# await self.connection_hook(conn)
|
|
||||||
|
|
||||||
async def connection_hook(self, conn: AsyncConnection):
|
|
||||||
info = await EnumInfo.fetch(conn, self.name)
|
|
||||||
if info is None:
|
|
||||||
raise ValueError(f"Enum {self.name} not found in database.")
|
|
||||||
register_enum(info, conn, self.enum, mapping=list(self.mapping.items()))
|
|
||||||
@@ -1,45 +0,0 @@
|
|||||||
from abc import abstractmethod
|
|
||||||
from typing import Any, Protocol, runtime_checkable
|
|
||||||
from itertools import chain
|
|
||||||
from psycopg import sql
|
|
||||||
|
|
||||||
|
|
||||||
@runtime_checkable
|
|
||||||
class Expression(Protocol):
|
|
||||||
__slots__ = ()
|
|
||||||
|
|
||||||
@abstractmethod
|
|
||||||
def as_tuple(self) -> tuple[sql.Composable, tuple[Any, ...]]:
|
|
||||||
raise NotImplementedError
|
|
||||||
|
|
||||||
|
|
||||||
class RawExpr(Expression):
|
|
||||||
__slots__ = ('expr', 'values')
|
|
||||||
|
|
||||||
expr: sql.Composable
|
|
||||||
values: tuple[Any, ...]
|
|
||||||
|
|
||||||
def __init__(self, expr: sql.Composable, values: tuple[Any, ...] = ()):
|
|
||||||
self.expr = expr
|
|
||||||
self.values = values
|
|
||||||
|
|
||||||
def as_tuple(self):
|
|
||||||
return (self.expr, self.values)
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def join(cls, *expressions: Expression, joiner: sql.SQL = sql.SQL(' ')):
|
|
||||||
"""
|
|
||||||
Join a sequence of Expressions into a single RawExpr.
|
|
||||||
"""
|
|
||||||
tups = (
|
|
||||||
expression.as_tuple()
|
|
||||||
for expression in expressions
|
|
||||||
)
|
|
||||||
return cls.join_tuples(*tups, joiner=joiner)
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def join_tuples(cls, *tuples: tuple[sql.Composable, tuple[Any, ...]], joiner: sql.SQL = sql.SQL(' ')):
|
|
||||||
exprs, values = zip(*tuples)
|
|
||||||
expr = joiner.join(exprs)
|
|
||||||
value = tuple(chain(*values))
|
|
||||||
return cls(expr, value)
|
|
||||||
@@ -1,155 +0,0 @@
|
|||||||
from typing import Any, Union, TypeVar, Generic, Type, overload, Optional, TYPE_CHECKING
|
|
||||||
from psycopg import sql
|
|
||||||
from datetime import datetime
|
|
||||||
|
|
||||||
from .base import RawExpr, Expression
|
|
||||||
from .conditions import Condition, Joiner
|
|
||||||
from .table import Table
|
|
||||||
|
|
||||||
|
|
||||||
class ColumnExpr(RawExpr):
|
|
||||||
__slots__ = ()
|
|
||||||
|
|
||||||
def __lt__(self, obj) -> Condition:
|
|
||||||
expr, values = self.as_tuple()
|
|
||||||
|
|
||||||
if isinstance(obj, Expression):
|
|
||||||
# column < Expression
|
|
||||||
obj_expr, obj_values = obj.as_tuple()
|
|
||||||
cond_exprs = (expr, Joiner.LT, obj_expr)
|
|
||||||
cond_values = (*values, *obj_values)
|
|
||||||
else:
|
|
||||||
# column < Literal
|
|
||||||
cond_exprs = (expr, Joiner.LT, sql.Placeholder())
|
|
||||||
cond_values = (*values, obj)
|
|
||||||
|
|
||||||
return Condition(cond_exprs[0], cond_exprs[1], cond_exprs[2], cond_values)
|
|
||||||
|
|
||||||
def __le__(self, obj) -> Condition:
|
|
||||||
expr, values = self.as_tuple()
|
|
||||||
|
|
||||||
if isinstance(obj, Expression):
|
|
||||||
# column <= Expression
|
|
||||||
obj_expr, obj_values = obj.as_tuple()
|
|
||||||
cond_exprs = (expr, Joiner.LE, obj_expr)
|
|
||||||
cond_values = (*values, *obj_values)
|
|
||||||
else:
|
|
||||||
# column <= Literal
|
|
||||||
cond_exprs = (expr, Joiner.LE, sql.Placeholder())
|
|
||||||
cond_values = (*values, obj)
|
|
||||||
|
|
||||||
return Condition(cond_exprs[0], cond_exprs[1], cond_exprs[2], cond_values)
|
|
||||||
|
|
||||||
def __eq__(self, obj) -> Condition: # type: ignore[override]
|
|
||||||
return Condition._expression_equality(self, obj)
|
|
||||||
|
|
||||||
def __ne__(self, obj) -> Condition: # type: ignore[override]
|
|
||||||
return ~(self.__eq__(obj))
|
|
||||||
|
|
||||||
def __gt__(self, obj) -> Condition:
|
|
||||||
return ~(self.__le__(obj))
|
|
||||||
|
|
||||||
def __ge__(self, obj) -> Condition:
|
|
||||||
return ~(self.__lt__(obj))
|
|
||||||
|
|
||||||
def __add__(self, obj: Union[Any, Expression]) -> 'ColumnExpr':
|
|
||||||
if isinstance(obj, Expression):
|
|
||||||
obj_expr, obj_values = obj.as_tuple()
|
|
||||||
return ColumnExpr(
|
|
||||||
sql.SQL("({} + {})").format(self.expr, obj_expr),
|
|
||||||
(*self.values, *obj_values)
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
return ColumnExpr(
|
|
||||||
sql.SQL("({} + {})").format(self.expr, sql.Placeholder()),
|
|
||||||
(*self.values, obj)
|
|
||||||
)
|
|
||||||
|
|
||||||
def __sub__(self, obj) -> 'ColumnExpr':
|
|
||||||
if isinstance(obj, Expression):
|
|
||||||
obj_expr, obj_values = obj.as_tuple()
|
|
||||||
return ColumnExpr(
|
|
||||||
sql.SQL("({} - {})").format(self.expr, obj_expr),
|
|
||||||
(*self.values, *obj_values)
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
return ColumnExpr(
|
|
||||||
sql.SQL("({} - {})").format(self.expr, sql.Placeholder()),
|
|
||||||
(*self.values, obj)
|
|
||||||
)
|
|
||||||
|
|
||||||
def __mul__(self, obj) -> 'ColumnExpr':
|
|
||||||
if isinstance(obj, Expression):
|
|
||||||
obj_expr, obj_values = obj.as_tuple()
|
|
||||||
return ColumnExpr(
|
|
||||||
sql.SQL("({} * {})").format(self.expr, obj_expr),
|
|
||||||
(*self.values, *obj_values)
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
return ColumnExpr(
|
|
||||||
sql.SQL("({} * {})").format(self.expr, sql.Placeholder()),
|
|
||||||
(*self.values, obj)
|
|
||||||
)
|
|
||||||
|
|
||||||
def CAST(self, target_type: sql.Composable):
|
|
||||||
return ColumnExpr(
|
|
||||||
sql.SQL("({}::{})").format(self.expr, target_type),
|
|
||||||
self.values
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
T = TypeVar('T')
|
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
|
||||||
from .models import RowModel
|
|
||||||
|
|
||||||
|
|
||||||
class Column(ColumnExpr, Generic[T]):
|
|
||||||
def __init__(self, name: Optional[str] = None,
|
|
||||||
primary: bool = False, references: Optional['Column'] = None,
|
|
||||||
type: Optional[Type[T]] = None):
|
|
||||||
self.primary = primary
|
|
||||||
self.references = references
|
|
||||||
self.name: str = name # type: ignore
|
|
||||||
self.owner: Optional['RowModel'] = None
|
|
||||||
self._type = type
|
|
||||||
|
|
||||||
self.expr = sql.Identifier(name) if name else sql.SQL('')
|
|
||||||
self.values = ()
|
|
||||||
|
|
||||||
def __set_name__(self, owner, name):
|
|
||||||
# Only allow setting the owner once
|
|
||||||
self.name = self.name or name
|
|
||||||
self.owner = owner
|
|
||||||
self.expr = sql.Identifier(self.owner._schema_, self.owner._tablename_, self.name)
|
|
||||||
|
|
||||||
@overload
|
|
||||||
def __get__(self: 'Column[T]', obj: None, objtype: "None | Type['RowModel']") -> 'Column[T]':
|
|
||||||
...
|
|
||||||
|
|
||||||
@overload
|
|
||||||
def __get__(self: 'Column[T]', obj: 'RowModel', objtype: Type['RowModel']) -> T:
|
|
||||||
...
|
|
||||||
|
|
||||||
def __get__(self: 'Column[T]', obj: "RowModel | None", objtype: "Type[RowModel] | None" = None) -> "T | Column[T]":
|
|
||||||
# Get value from row data or session
|
|
||||||
if obj is None:
|
|
||||||
return self
|
|
||||||
else:
|
|
||||||
return obj.data[self.name]
|
|
||||||
|
|
||||||
|
|
||||||
class Integer(Column[int]):
|
|
||||||
pass
|
|
||||||
|
|
||||||
|
|
||||||
class String(Column[str]):
|
|
||||||
pass
|
|
||||||
|
|
||||||
|
|
||||||
class Bool(Column[bool]):
|
|
||||||
pass
|
|
||||||
|
|
||||||
|
|
||||||
class Timestamp(Column[datetime]):
|
|
||||||
pass
|
|
||||||
@@ -1,214 +0,0 @@
|
|||||||
# from meta import sharding
|
|
||||||
from typing import Any, Union
|
|
||||||
from enum import Enum
|
|
||||||
from itertools import chain
|
|
||||||
from psycopg import sql
|
|
||||||
|
|
||||||
from .base import Expression, RawExpr
|
|
||||||
|
|
||||||
|
|
||||||
"""
|
|
||||||
A Condition is a "logical" database expression, intended for use in Where statements.
|
|
||||||
Conditions support bitwise logical operators ~, &, |, each producing another Condition.
|
|
||||||
"""
|
|
||||||
|
|
||||||
NULL = None
|
|
||||||
|
|
||||||
|
|
||||||
class Joiner(Enum):
|
|
||||||
EQUALS = ('=', '!=')
|
|
||||||
IS = ('IS', 'IS NOT')
|
|
||||||
LIKE = ('LIKE', 'NOT LIKE')
|
|
||||||
BETWEEN = ('BETWEEN', 'NOT BETWEEN')
|
|
||||||
IN = ('IN', 'NOT IN')
|
|
||||||
LT = ('<', '>=')
|
|
||||||
LE = ('<=', '>')
|
|
||||||
NONE = ('', '')
|
|
||||||
|
|
||||||
|
|
||||||
class Condition(Expression):
|
|
||||||
__slots__ = ('expr1', 'joiner', 'negated', 'expr2', 'values')
|
|
||||||
|
|
||||||
def __init__(self,
|
|
||||||
expr1: sql.Composable, joiner: Joiner = Joiner.NONE, expr2: sql.Composable = sql.SQL(''),
|
|
||||||
values: tuple[Any, ...] = (), negated=False
|
|
||||||
):
|
|
||||||
self.expr1 = expr1
|
|
||||||
self.joiner = joiner
|
|
||||||
self.negated = negated
|
|
||||||
self.expr2 = expr2
|
|
||||||
self.values = values
|
|
||||||
|
|
||||||
def as_tuple(self):
|
|
||||||
expr = sql.SQL(' ').join((self.expr1, sql.SQL(self.joiner.value[self.negated]), self.expr2))
|
|
||||||
if self.negated and self.joiner is Joiner.NONE:
|
|
||||||
expr = sql.SQL("NOT ({})").format(expr)
|
|
||||||
return (expr, self.values)
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def construct(cls, *conditions: 'Condition', **kwargs: Union[Any, Expression]):
|
|
||||||
"""
|
|
||||||
Construct a Condition from a sequence of Conditions,
|
|
||||||
together with some explicit column conditions.
|
|
||||||
"""
|
|
||||||
# TODO: Consider adding a _table identifier here so we can identify implicit columns
|
|
||||||
# Or just require subquery type conditions to always come from modelled tables.
|
|
||||||
implicit_conditions = (
|
|
||||||
cls._expression_equality(RawExpr(sql.Identifier(column)), value) for column, value in kwargs.items()
|
|
||||||
)
|
|
||||||
return cls._and(*conditions, *implicit_conditions)
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def _and(cls, *conditions: 'Condition'):
|
|
||||||
if not len(conditions):
|
|
||||||
raise ValueError("Cannot combine 0 Conditions")
|
|
||||||
if len(conditions) == 1:
|
|
||||||
return conditions[0]
|
|
||||||
|
|
||||||
exprs, values = zip(*(condition.as_tuple() for condition in conditions))
|
|
||||||
cond_expr = sql.SQL(' AND ').join((sql.SQL('({})').format(expr) for expr in exprs))
|
|
||||||
cond_values = tuple(chain(*values))
|
|
||||||
|
|
||||||
return Condition(cond_expr, values=cond_values)
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def _or(cls, *conditions: 'Condition'):
|
|
||||||
if not len(conditions):
|
|
||||||
raise ValueError("Cannot combine 0 Conditions")
|
|
||||||
if len(conditions) == 1:
|
|
||||||
return conditions[0]
|
|
||||||
|
|
||||||
exprs, values = zip(*(condition.as_tuple() for condition in conditions))
|
|
||||||
cond_expr = sql.SQL(' OR ').join((sql.SQL('({})').format(expr) for expr in exprs))
|
|
||||||
cond_values = tuple(chain(*values))
|
|
||||||
|
|
||||||
return Condition(cond_expr, values=cond_values)
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def _not(cls, condition: 'Condition'):
|
|
||||||
condition.negated = not condition.negated
|
|
||||||
return condition
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def _expression_equality(cls, column: Expression, value: Union[Any, Expression]) -> 'Condition':
|
|
||||||
# TODO: Check if this supports sbqueries
|
|
||||||
col_expr, col_values = column.as_tuple()
|
|
||||||
|
|
||||||
# TODO: Also support sql.SQL? For joins?
|
|
||||||
if isinstance(value, Expression):
|
|
||||||
# column = Expression
|
|
||||||
value_expr, value_values = value.as_tuple()
|
|
||||||
cond_exprs = (col_expr, Joiner.EQUALS, value_expr)
|
|
||||||
cond_values = (*col_values, *value_values)
|
|
||||||
elif isinstance(value, (tuple, list)):
|
|
||||||
# column in (...)
|
|
||||||
# TODO: Support expressions in value tuple?
|
|
||||||
if not value:
|
|
||||||
raise ValueError("Cannot create Condition from empty iterable!")
|
|
||||||
value_expr = sql.SQL('({})').format(sql.SQL(',').join(sql.Placeholder() * len(value)))
|
|
||||||
cond_exprs = (col_expr, Joiner.IN, value_expr)
|
|
||||||
cond_values = (*col_values, *value)
|
|
||||||
elif value is None:
|
|
||||||
# column IS NULL
|
|
||||||
cond_exprs = (col_expr, Joiner.IS, sql.NULL)
|
|
||||||
cond_values = col_values
|
|
||||||
else:
|
|
||||||
# column = Literal
|
|
||||||
cond_exprs = (col_expr, Joiner.EQUALS, sql.Placeholder())
|
|
||||||
cond_values = (*col_values, value)
|
|
||||||
|
|
||||||
return cls(cond_exprs[0], cond_exprs[1], cond_exprs[2], cond_values)
|
|
||||||
|
|
||||||
def __invert__(self) -> 'Condition':
|
|
||||||
self.negated = not self.negated
|
|
||||||
return self
|
|
||||||
|
|
||||||
def __and__(self, condition: 'Condition') -> 'Condition':
|
|
||||||
return self._and(self, condition)
|
|
||||||
|
|
||||||
def __or__(self, condition: 'Condition') -> 'Condition':
|
|
||||||
return self._or(self, condition)
|
|
||||||
|
|
||||||
|
|
||||||
# Helper method to simply condition construction
|
|
||||||
def condition(*args, **kwargs) -> Condition:
|
|
||||||
return Condition.construct(*args, **kwargs)
|
|
||||||
|
|
||||||
|
|
||||||
# class NOT(Condition):
|
|
||||||
# __slots__ = ('value',)
|
|
||||||
#
|
|
||||||
# def __init__(self, value):
|
|
||||||
# self.value = value
|
|
||||||
#
|
|
||||||
# def apply(self, key, values, conditions):
|
|
||||||
# item = self.value
|
|
||||||
# if isinstance(item, (list, tuple)):
|
|
||||||
# if item:
|
|
||||||
# conditions.append("{} NOT IN ({})".format(key, ", ".join([_replace_char] * len(item))))
|
|
||||||
# values.extend(item)
|
|
||||||
# else:
|
|
||||||
# raise ValueError("Cannot check an empty iterable!")
|
|
||||||
# else:
|
|
||||||
# conditions.append("{}!={}".format(key, _replace_char))
|
|
||||||
# values.append(item)
|
|
||||||
#
|
|
||||||
#
|
|
||||||
# class GEQ(Condition):
|
|
||||||
# __slots__ = ('value',)
|
|
||||||
#
|
|
||||||
# def __init__(self, value):
|
|
||||||
# self.value = value
|
|
||||||
#
|
|
||||||
# def apply(self, key, values, conditions):
|
|
||||||
# item = self.value
|
|
||||||
# if isinstance(item, (list, tuple)):
|
|
||||||
# raise ValueError("Cannot apply GEQ condition to a list!")
|
|
||||||
# else:
|
|
||||||
# conditions.append("{} >= {}".format(key, _replace_char))
|
|
||||||
# values.append(item)
|
|
||||||
#
|
|
||||||
#
|
|
||||||
# class LEQ(Condition):
|
|
||||||
# __slots__ = ('value',)
|
|
||||||
#
|
|
||||||
# def __init__(self, value):
|
|
||||||
# self.value = value
|
|
||||||
#
|
|
||||||
# def apply(self, key, values, conditions):
|
|
||||||
# item = self.value
|
|
||||||
# if isinstance(item, (list, tuple)):
|
|
||||||
# raise ValueError("Cannot apply LEQ condition to a list!")
|
|
||||||
# else:
|
|
||||||
# conditions.append("{} <= {}".format(key, _replace_char))
|
|
||||||
# values.append(item)
|
|
||||||
#
|
|
||||||
#
|
|
||||||
# class Constant(Condition):
|
|
||||||
# __slots__ = ('value',)
|
|
||||||
#
|
|
||||||
# def __init__(self, value):
|
|
||||||
# self.value = value
|
|
||||||
#
|
|
||||||
# def apply(self, key, values, conditions):
|
|
||||||
# conditions.append("{} {}".format(key, self.value))
|
|
||||||
#
|
|
||||||
#
|
|
||||||
# class SHARDID(Condition):
|
|
||||||
# __slots__ = ('shardid', 'shard_count')
|
|
||||||
#
|
|
||||||
# def __init__(self, shardid, shard_count):
|
|
||||||
# self.shardid = shardid
|
|
||||||
# self.shard_count = shard_count
|
|
||||||
#
|
|
||||||
# def apply(self, key, values, conditions):
|
|
||||||
# if self.shard_count > 1:
|
|
||||||
# conditions.append("({} >> 22) %% {} = {}".format(key, self.shard_count, _replace_char))
|
|
||||||
# values.append(self.shardid)
|
|
||||||
#
|
|
||||||
#
|
|
||||||
# # THIS_SHARD = SHARDID(sharding.shard_number, sharding.shard_count)
|
|
||||||
#
|
|
||||||
#
|
|
||||||
# NULL = Constant('IS NULL')
|
|
||||||
# NOTNULL = Constant('IS NOT NULL')
|
|
||||||
@@ -1,135 +0,0 @@
|
|||||||
from typing import Protocol, runtime_checkable, Callable, Awaitable, Optional
|
|
||||||
import logging
|
|
||||||
|
|
||||||
from contextvars import ContextVar
|
|
||||||
from contextlib import asynccontextmanager
|
|
||||||
import psycopg as psq
|
|
||||||
from psycopg_pool import AsyncConnectionPool
|
|
||||||
from psycopg.pq import TransactionStatus
|
|
||||||
|
|
||||||
from .cursor import AsyncLoggingCursor
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
|
||||||
|
|
||||||
row_factory = psq.rows.dict_row
|
|
||||||
|
|
||||||
ctx_connection: Optional[ContextVar[psq.AsyncConnection]] = ContextVar('connection', default=None)
|
|
||||||
|
|
||||||
|
|
||||||
class Connector:
|
|
||||||
cursor_factory = AsyncLoggingCursor
|
|
||||||
|
|
||||||
def __init__(self, conn_args):
|
|
||||||
self._conn_args = conn_args
|
|
||||||
self._conn_kwargs = dict(autocommit=True, row_factory=row_factory, cursor_factory=self.cursor_factory)
|
|
||||||
|
|
||||||
self.pool = self.make_pool()
|
|
||||||
|
|
||||||
self.conn_hooks = []
|
|
||||||
|
|
||||||
@property
|
|
||||||
def conn(self) -> Optional[psq.AsyncConnection]:
|
|
||||||
"""
|
|
||||||
Convenience property for the current context connection.
|
|
||||||
"""
|
|
||||||
return ctx_connection.get()
|
|
||||||
|
|
||||||
@conn.setter
|
|
||||||
def conn(self, conn: psq.AsyncConnection):
|
|
||||||
"""
|
|
||||||
Set the contextual connection in the current context.
|
|
||||||
Always do this in an isolated context!
|
|
||||||
"""
|
|
||||||
ctx_connection.set(conn)
|
|
||||||
|
|
||||||
def make_pool(self) -> AsyncConnectionPool:
|
|
||||||
logger.info("Initialising connection pool.", extra={'action': "Pool Init"})
|
|
||||||
return AsyncConnectionPool(
|
|
||||||
self._conn_args,
|
|
||||||
open=False,
|
|
||||||
min_size=4,
|
|
||||||
max_size=8,
|
|
||||||
configure=self._setup_connection,
|
|
||||||
kwargs=self._conn_kwargs
|
|
||||||
)
|
|
||||||
|
|
||||||
async def refresh_pool(self):
|
|
||||||
"""
|
|
||||||
Refresh the pool.
|
|
||||||
|
|
||||||
The point of this is to invalidate any existing connections so that the connection set up is run again.
|
|
||||||
Better ways should be sought (a way to
|
|
||||||
"""
|
|
||||||
logger.info("Pool refresh requested, closing and reopening.")
|
|
||||||
old_pool = self.pool
|
|
||||||
self.pool = self.make_pool()
|
|
||||||
await self.pool.open()
|
|
||||||
logger.info(f"Old pool statistics: {self.pool.get_stats()}")
|
|
||||||
await old_pool.close()
|
|
||||||
logger.info("Pool refresh complete.")
|
|
||||||
|
|
||||||
async def map_over_pool(self, callable):
|
|
||||||
"""
|
|
||||||
Dangerous method to call a method on each connection in the pool.
|
|
||||||
|
|
||||||
Utilises private methods of the AsyncConnectionPool.
|
|
||||||
"""
|
|
||||||
async with self.pool._lock:
|
|
||||||
conns = list(self.pool._pool)
|
|
||||||
while conns:
|
|
||||||
conn = conns.pop()
|
|
||||||
try:
|
|
||||||
await callable(conn)
|
|
||||||
except Exception:
|
|
||||||
logger.exception(f"Mapped connection task failed. {callable.__name__}")
|
|
||||||
|
|
||||||
@asynccontextmanager
|
|
||||||
async def open(self):
|
|
||||||
try:
|
|
||||||
logger.info("Opening database pool.")
|
|
||||||
await self.pool.open()
|
|
||||||
yield
|
|
||||||
finally:
|
|
||||||
# May be a different pool!
|
|
||||||
logger.info(f"Closing database pool. Pool statistics: {self.pool.get_stats()}")
|
|
||||||
await self.pool.close()
|
|
||||||
|
|
||||||
@asynccontextmanager
|
|
||||||
async def connection(self) -> psq.AsyncConnection:
|
|
||||||
"""
|
|
||||||
Asynchronous context manager to get and manage a connection.
|
|
||||||
|
|
||||||
If the context connection is set, uses this and does not manage the lifetime.
|
|
||||||
Otherwise, requests a new connection from the pool and returns it when done.
|
|
||||||
"""
|
|
||||||
logger.debug("Database connection requested.", extra={'action': "Data Connect"})
|
|
||||||
if (conn := self.conn):
|
|
||||||
yield conn
|
|
||||||
else:
|
|
||||||
async with self.pool.connection() as conn:
|
|
||||||
yield conn
|
|
||||||
|
|
||||||
async def _setup_connection(self, conn: psq.AsyncConnection):
|
|
||||||
logger.debug("Initialising new connection.", extra={'action': "Conn Init"})
|
|
||||||
for hook in self.conn_hooks:
|
|
||||||
try:
|
|
||||||
await hook(conn)
|
|
||||||
except Exception:
|
|
||||||
logger.exception("Exception encountered setting up new connection")
|
|
||||||
return conn
|
|
||||||
|
|
||||||
def connect_hook(self, coro: Callable[[psq.AsyncConnection], Awaitable[None]]):
|
|
||||||
"""
|
|
||||||
Minimal decorator to register a coroutine to run on connect or reconnect.
|
|
||||||
|
|
||||||
Note that these are only run on connect and reconnect.
|
|
||||||
If a hook is registered after connection, it will not be run.
|
|
||||||
"""
|
|
||||||
self.conn_hooks.append(coro)
|
|
||||||
return coro
|
|
||||||
|
|
||||||
|
|
||||||
@runtime_checkable
|
|
||||||
class Connectable(Protocol):
|
|
||||||
def bind(self, connector: Connector):
|
|
||||||
raise NotImplementedError
|
|
||||||
@@ -1,42 +0,0 @@
|
|||||||
import logging
|
|
||||||
from typing import Optional
|
|
||||||
|
|
||||||
from psycopg import AsyncCursor, sql
|
|
||||||
from psycopg.abc import Query, Params
|
|
||||||
from psycopg._encodings import pgconn_encoding
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
|
||||||
|
|
||||||
|
|
||||||
class AsyncLoggingCursor(AsyncCursor):
|
|
||||||
def mogrify_query(self, query: Query):
|
|
||||||
if isinstance(query, str):
|
|
||||||
msg = query
|
|
||||||
elif isinstance(query, (sql.SQL, sql.Composed)):
|
|
||||||
msg = query.as_string(self)
|
|
||||||
elif isinstance(query, bytes):
|
|
||||||
msg = query.decode(pgconn_encoding(self._conn.pgconn), 'replace')
|
|
||||||
else:
|
|
||||||
msg = repr(query)
|
|
||||||
return msg
|
|
||||||
|
|
||||||
async def execute(self, query: Query, params: Optional[Params] = None, **kwargs):
|
|
||||||
if logging.DEBUG >= logger.getEffectiveLevel():
|
|
||||||
msg = self.mogrify_query(query)
|
|
||||||
logger.debug(
|
|
||||||
"Executing query (%s) with values %s", msg, params,
|
|
||||||
extra={'action': "Query Execute"}
|
|
||||||
)
|
|
||||||
try:
|
|
||||||
return await super().execute(query, params=params, **kwargs)
|
|
||||||
except Exception:
|
|
||||||
msg = self.mogrify_query(query)
|
|
||||||
logger.exception(
|
|
||||||
"Exception during query execution. Query (%s) with parameters %s.",
|
|
||||||
msg, params,
|
|
||||||
extra={'action': "Query Execute"},
|
|
||||||
stack_info=True
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
# TODO: Possibly log execution time
|
|
||||||
pass
|
|
||||||
@@ -1,47 +0,0 @@
|
|||||||
from typing import TypeVar
|
|
||||||
import logging
|
|
||||||
from collections import namedtuple
|
|
||||||
|
|
||||||
# from .cursor import AsyncLoggingCursor
|
|
||||||
from .registry import Registry
|
|
||||||
from .connector import Connector
|
|
||||||
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
|
||||||
|
|
||||||
Version = namedtuple('Version', ('version', 'time', 'author'))
|
|
||||||
|
|
||||||
T = TypeVar('T', bound=Registry)
|
|
||||||
|
|
||||||
|
|
||||||
class Database(Connector):
|
|
||||||
# cursor_factory = AsyncLoggingCursor
|
|
||||||
|
|
||||||
def __init__(self, *args, **kwargs):
|
|
||||||
super().__init__(*args, **kwargs)
|
|
||||||
|
|
||||||
self.registries: dict[str, Registry] = {}
|
|
||||||
|
|
||||||
def load_registry(self, registry: T) -> T:
|
|
||||||
logger.debug(
|
|
||||||
f"Loading and binding registry '{registry.name}'.",
|
|
||||||
extra={'action': f"Reg {registry.name}"}
|
|
||||||
)
|
|
||||||
registry.bind(self)
|
|
||||||
self.registries[registry.name] = registry
|
|
||||||
return registry
|
|
||||||
|
|
||||||
async def version(self) -> Version:
|
|
||||||
"""
|
|
||||||
Return the current schema version as a Version namedtuple.
|
|
||||||
"""
|
|
||||||
async with self.connection() as conn:
|
|
||||||
async with conn.cursor() as cursor:
|
|
||||||
# Get last entry in version table, compare against desired version
|
|
||||||
await cursor.execute("SELECT * FROM VersionHistory ORDER BY time DESC LIMIT 1")
|
|
||||||
row = await cursor.fetchone()
|
|
||||||
if row:
|
|
||||||
return Version(row['version'], row['time'], row['author'])
|
|
||||||
else:
|
|
||||||
# No versions in the database
|
|
||||||
return Version(-1, None, None)
|
|
||||||
@@ -1,323 +0,0 @@
|
|||||||
from typing import TypeVar, Type, Optional, Generic, Union
|
|
||||||
# from typing_extensions import Self
|
|
||||||
from weakref import WeakValueDictionary
|
|
||||||
from collections.abc import MutableMapping
|
|
||||||
|
|
||||||
from psycopg.rows import DictRow
|
|
||||||
|
|
||||||
from .table import Table
|
|
||||||
from .columns import Column
|
|
||||||
from . import queries as q
|
|
||||||
from .connector import Connector
|
|
||||||
from .registry import Registry
|
|
||||||
|
|
||||||
|
|
||||||
RowT = TypeVar('RowT', bound='RowModel')
|
|
||||||
|
|
||||||
|
|
||||||
class MISSING:
|
|
||||||
__slots__ = ('oid',)
|
|
||||||
|
|
||||||
def __init__(self, oid):
|
|
||||||
self.oid = oid
|
|
||||||
|
|
||||||
|
|
||||||
class RowTable(Table, Generic[RowT]):
|
|
||||||
__slots__ = (
|
|
||||||
'model',
|
|
||||||
)
|
|
||||||
|
|
||||||
def __init__(self, name, model: Type[RowT], **kwargs):
|
|
||||||
super().__init__(name, **kwargs)
|
|
||||||
self.model = model
|
|
||||||
|
|
||||||
@property
|
|
||||||
def columns(self):
|
|
||||||
return self.model._columns_
|
|
||||||
|
|
||||||
@property
|
|
||||||
def id_col(self):
|
|
||||||
return self.model._key_
|
|
||||||
|
|
||||||
@property
|
|
||||||
def row_cache(self):
|
|
||||||
return self.model._cache_
|
|
||||||
|
|
||||||
def _many_query_adapter(self, *data):
|
|
||||||
self.model._make_rows(*data)
|
|
||||||
return data
|
|
||||||
|
|
||||||
def _single_query_adapter(self, *data):
|
|
||||||
if data:
|
|
||||||
self.model._make_rows(*data)
|
|
||||||
return data[0]
|
|
||||||
else:
|
|
||||||
return None
|
|
||||||
|
|
||||||
def _delete_query_adapter(self, *data):
|
|
||||||
self.model._delete_rows(*data)
|
|
||||||
return data
|
|
||||||
|
|
||||||
# New methods to fetch and create rows
|
|
||||||
async def create_row(self, *args, **kwargs) -> RowT:
|
|
||||||
data = await super().insert(*args, **kwargs)
|
|
||||||
return self.model._make_rows(data)[0]
|
|
||||||
|
|
||||||
def fetch_rows_where(self, *args, **kwargs) -> q.Select[list[RowT]]:
|
|
||||||
# TODO: Handle list of rowids here?
|
|
||||||
return q.Select(
|
|
||||||
self.identifier,
|
|
||||||
row_adapter=self.model._make_rows,
|
|
||||||
connector=self.connector
|
|
||||||
).where(*args, **kwargs)
|
|
||||||
|
|
||||||
|
|
||||||
WK = TypeVar('WK')
|
|
||||||
WV = TypeVar('WV')
|
|
||||||
|
|
||||||
|
|
||||||
class WeakCache(Generic[WK, WV], MutableMapping[WK, WV]):
|
|
||||||
def __init__(self, ref_cache):
|
|
||||||
self.ref_cache = ref_cache
|
|
||||||
self.weak_cache = WeakValueDictionary()
|
|
||||||
|
|
||||||
def __getitem__(self, key):
|
|
||||||
value = self.weak_cache[key]
|
|
||||||
self.ref_cache[key] = value
|
|
||||||
return value
|
|
||||||
|
|
||||||
def __setitem__(self, key, value):
|
|
||||||
self.weak_cache[key] = value
|
|
||||||
self.ref_cache[key] = value
|
|
||||||
|
|
||||||
def __delitem__(self, key):
|
|
||||||
del self.weak_cache[key]
|
|
||||||
try:
|
|
||||||
del self.ref_cache[key]
|
|
||||||
except KeyError:
|
|
||||||
pass
|
|
||||||
|
|
||||||
def __contains__(self, key):
|
|
||||||
return key in self.weak_cache
|
|
||||||
|
|
||||||
def __iter__(self):
|
|
||||||
return iter(self.weak_cache)
|
|
||||||
|
|
||||||
def __len__(self):
|
|
||||||
return len(self.weak_cache)
|
|
||||||
|
|
||||||
def get(self, key, default=None):
|
|
||||||
try:
|
|
||||||
return self[key]
|
|
||||||
except KeyError:
|
|
||||||
return default
|
|
||||||
|
|
||||||
def pop(self, key, default=None):
|
|
||||||
if key in self:
|
|
||||||
value = self[key]
|
|
||||||
del self[key]
|
|
||||||
else:
|
|
||||||
value = default
|
|
||||||
return value
|
|
||||||
|
|
||||||
|
|
||||||
# TODO: Implement getitem and setitem, for dynamic column access
|
|
||||||
class RowModel:
|
|
||||||
__slots__ = ('data',)
|
|
||||||
|
|
||||||
_schema_: str = 'public'
|
|
||||||
_tablename_: Optional[str] = None
|
|
||||||
_columns_: dict[str, Column] = {}
|
|
||||||
|
|
||||||
# Cache to keep track of registered Rows
|
|
||||||
_cache_: Union[dict, WeakValueDictionary, WeakCache] = None # type: ignore
|
|
||||||
|
|
||||||
_key_: tuple[str, ...] = ()
|
|
||||||
_connector: Optional[Connector] = None
|
|
||||||
_registry: Optional[Registry] = None
|
|
||||||
|
|
||||||
# TODO: Proper typing for a classvariable which gets dynamically assigned in subclass
|
|
||||||
table: RowTable = None
|
|
||||||
|
|
||||||
def __init_subclass__(cls: Type[RowT], table: Optional[str] = None):
|
|
||||||
"""
|
|
||||||
Set table, _columns_, and _key_.
|
|
||||||
"""
|
|
||||||
if table is not None:
|
|
||||||
cls._tablename_ = table
|
|
||||||
|
|
||||||
if cls._tablename_ is not None:
|
|
||||||
columns = {}
|
|
||||||
for key, value in cls.__dict__.items():
|
|
||||||
if isinstance(value, Column):
|
|
||||||
columns[key] = value
|
|
||||||
|
|
||||||
cls._columns_ = columns
|
|
||||||
if not cls._key_:
|
|
||||||
cls._key_ = tuple(column.name for column in columns.values() if column.primary)
|
|
||||||
cls.table = RowTable(cls._tablename_, cls, schema=cls._schema_)
|
|
||||||
if cls._cache_ is None:
|
|
||||||
cls._cache_ = WeakValueDictionary()
|
|
||||||
|
|
||||||
def __new__(cls, data):
|
|
||||||
# Registry pattern.
|
|
||||||
# Ensure each rowid always refers to a single Model instance
|
|
||||||
if data is not None:
|
|
||||||
rowid = cls._id_from_data(data)
|
|
||||||
|
|
||||||
cache = cls._cache_
|
|
||||||
|
|
||||||
if (row := cache.get(rowid, None)) is not None:
|
|
||||||
obj = row
|
|
||||||
else:
|
|
||||||
obj = cache[rowid] = super().__new__(cls)
|
|
||||||
else:
|
|
||||||
obj = super().__new__(cls)
|
|
||||||
|
|
||||||
return obj
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def as_tuple(cls):
|
|
||||||
return (cls.table.identifier, ())
|
|
||||||
|
|
||||||
def __init__(self, data):
|
|
||||||
self.data = data
|
|
||||||
|
|
||||||
def __getitem__(self, key):
|
|
||||||
return self.data[key]
|
|
||||||
|
|
||||||
def __setitem__(self, key, value):
|
|
||||||
self.data[key] = value
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def bind(cls, connector: Connector):
|
|
||||||
if cls.table is None:
|
|
||||||
raise ValueError("Cannot bind abstract RowModel")
|
|
||||||
cls._connector = connector
|
|
||||||
cls.table.bind(connector)
|
|
||||||
return cls
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def attach_to(cls, registry: Registry):
|
|
||||||
cls._registry = registry
|
|
||||||
return cls
|
|
||||||
|
|
||||||
@property
|
|
||||||
def _dict_(self):
|
|
||||||
return {key: self.data[key] for key in self._key_}
|
|
||||||
|
|
||||||
@property
|
|
||||||
def _rowid_(self):
|
|
||||||
return tuple(self.data[key] for key in self._key_)
|
|
||||||
|
|
||||||
def __repr__(self):
|
|
||||||
return "{}.{}({})".format(
|
|
||||||
self.table.schema,
|
|
||||||
self.table.name,
|
|
||||||
', '.join(repr(column.__get__(self)) for column in self._columns_.values())
|
|
||||||
)
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def _id_from_data(cls, data):
|
|
||||||
return tuple(data[key] for key in cls._key_)
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def _dict_from_id(cls, rowid):
|
|
||||||
return dict(zip(cls._key_, rowid))
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def _make_rows(cls: Type[RowT], *data_rows: DictRow) -> list[RowT]:
|
|
||||||
"""
|
|
||||||
Create or retrieve Row objects for each provided data row.
|
|
||||||
If the rows already exist in cache, updates the cached row.
|
|
||||||
"""
|
|
||||||
# TODO: Handle partial row data here somehow?
|
|
||||||
rows = [cls(data_row) for data_row in data_rows]
|
|
||||||
return rows
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def _delete_rows(cls, *data_rows):
|
|
||||||
"""
|
|
||||||
Remove the given rows from cache, if they exist.
|
|
||||||
May be extended to handle object deletion.
|
|
||||||
"""
|
|
||||||
cache = cls._cache_
|
|
||||||
|
|
||||||
for data_row in data_rows:
|
|
||||||
rowid = cls._id_from_data(data_row)
|
|
||||||
cache.pop(rowid, None)
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
async def create(cls: Type[RowT], *args, **kwargs) -> RowT:
|
|
||||||
return await cls.table.create_row(*args, **kwargs)
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def fetch_where(cls: Type[RowT], *args, **kwargs):
|
|
||||||
return cls.table.fetch_rows_where(*args, **kwargs)
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
async def fetch(cls: Type[RowT], *rowid, cached=True) -> Optional[RowT]:
|
|
||||||
"""
|
|
||||||
Fetch the row with the given id, retrieving from cache where possible.
|
|
||||||
"""
|
|
||||||
row = cls._cache_.get(rowid, None) if cached else None
|
|
||||||
if row is None:
|
|
||||||
rows = await cls.fetch_where(**cls._dict_from_id(rowid))
|
|
||||||
row = rows[0] if rows else None
|
|
||||||
if row is None:
|
|
||||||
cls._cache_[rowid] = cls(None)
|
|
||||||
elif row.data is None:
|
|
||||||
row = None
|
|
||||||
|
|
||||||
return row
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
async def fetch_or_create(cls, *rowid, **kwargs):
|
|
||||||
"""
|
|
||||||
Helper method to fetch a row with the given id or fields, or create it if it doesn't exist.
|
|
||||||
"""
|
|
||||||
if rowid:
|
|
||||||
row = await cls.fetch(*rowid)
|
|
||||||
else:
|
|
||||||
rows = await cls.fetch_where(**kwargs).limit(1)
|
|
||||||
row = rows[0] if rows else None
|
|
||||||
|
|
||||||
if row is None:
|
|
||||||
creation_kwargs = kwargs
|
|
||||||
if rowid:
|
|
||||||
creation_kwargs.update(cls._dict_from_id(rowid))
|
|
||||||
row = await cls.create(**creation_kwargs)
|
|
||||||
return row
|
|
||||||
|
|
||||||
async def refresh(self: RowT) -> Optional[RowT]:
|
|
||||||
"""
|
|
||||||
Refresh this Row from data.
|
|
||||||
|
|
||||||
The return value may be `None` if the row was deleted.
|
|
||||||
"""
|
|
||||||
rows = await self.table.select_where(**self._dict_)
|
|
||||||
if not rows:
|
|
||||||
return None
|
|
||||||
else:
|
|
||||||
self.data = rows[0]
|
|
||||||
return self
|
|
||||||
|
|
||||||
async def update(self: RowT, **values) -> Optional[RowT]:
|
|
||||||
"""
|
|
||||||
Update this Row with the given values.
|
|
||||||
|
|
||||||
Internally passes the provided `values` to the `update` Query.
|
|
||||||
The return value may be `None` if the row was deleted.
|
|
||||||
"""
|
|
||||||
data = await self.table.update_where(**self._dict_).set(**values).with_adapter(self._make_rows)
|
|
||||||
if not data:
|
|
||||||
return None
|
|
||||||
else:
|
|
||||||
return data[0]
|
|
||||||
|
|
||||||
async def delete(self: RowT) -> Optional[RowT]:
|
|
||||||
"""
|
|
||||||
Delete this Row.
|
|
||||||
"""
|
|
||||||
data = await self.table.delete_where(**self._dict_).with_adapter(self._delete_rows)
|
|
||||||
return data[0] if data is not None else None
|
|
||||||
@@ -1,644 +0,0 @@
|
|||||||
from typing import Optional, TypeVar, Any, Callable, Generic, List, Union
|
|
||||||
from enum import Enum
|
|
||||||
from itertools import chain
|
|
||||||
from psycopg import AsyncConnection, AsyncCursor
|
|
||||||
from psycopg import sql
|
|
||||||
from psycopg.rows import DictRow
|
|
||||||
|
|
||||||
import logging
|
|
||||||
|
|
||||||
from .conditions import Condition
|
|
||||||
from .base import Expression, RawExpr
|
|
||||||
from .connector import Connector
|
|
||||||
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
|
||||||
|
|
||||||
|
|
||||||
TQueryT = TypeVar('TQueryT', bound='TableQuery')
|
|
||||||
SQueryT = TypeVar('SQueryT', bound='Select')
|
|
||||||
|
|
||||||
QueryResult = TypeVar('QueryResult')
|
|
||||||
|
|
||||||
|
|
||||||
class Query(Generic[QueryResult]):
|
|
||||||
"""
|
|
||||||
ABC for an executable query statement.
|
|
||||||
"""
|
|
||||||
__slots__ = ('conn', 'cursor', '_adapter', 'connector', 'result')
|
|
||||||
|
|
||||||
_adapter: Callable[..., QueryResult]
|
|
||||||
|
|
||||||
def __init__(self, *args, row_adapter=None, connector=None, conn=None, cursor=None, **kwargs):
|
|
||||||
self.connector: Optional[Connector] = connector
|
|
||||||
self.conn: Optional[AsyncConnection] = conn
|
|
||||||
self.cursor: Optional[AsyncCursor] = cursor
|
|
||||||
|
|
||||||
if row_adapter is not None:
|
|
||||||
self._adapter = row_adapter
|
|
||||||
else:
|
|
||||||
self._adapter = self._no_adapter
|
|
||||||
|
|
||||||
self.result: Optional[QueryResult] = None
|
|
||||||
|
|
||||||
def bind(self, connector: Connector):
|
|
||||||
self.connector = connector
|
|
||||||
return self
|
|
||||||
|
|
||||||
def with_cursor(self, cursor: AsyncCursor):
|
|
||||||
self.cursor = cursor
|
|
||||||
return self
|
|
||||||
|
|
||||||
def with_connection(self, conn: AsyncConnection):
|
|
||||||
self.conn = conn
|
|
||||||
return self
|
|
||||||
|
|
||||||
def _no_adapter(self, *data: DictRow) -> tuple[DictRow, ...]:
|
|
||||||
return data
|
|
||||||
|
|
||||||
def with_adapter(self, callable: Callable[..., QueryResult]):
|
|
||||||
# NOTE: Postcomposition functor, Query[QR2] = (QR1 -> QR2) o Query[QR1]
|
|
||||||
# For this to work cleanly, callable should have arg type of QR1, not any
|
|
||||||
self._adapter = callable
|
|
||||||
return self
|
|
||||||
|
|
||||||
def with_no_adapter(self):
|
|
||||||
"""
|
|
||||||
Sets the adapater to the identity.
|
|
||||||
"""
|
|
||||||
self._adapter = self._no_adapter
|
|
||||||
return self
|
|
||||||
|
|
||||||
def one(self):
|
|
||||||
# TODO: Postcomposition with item functor, Query[List[QR1]] -> Query[QR1]
|
|
||||||
return self
|
|
||||||
|
|
||||||
def build(self) -> Expression:
|
|
||||||
raise NotImplementedError
|
|
||||||
|
|
||||||
async def _execute(self, cursor: AsyncCursor) -> QueryResult:
|
|
||||||
query, values = self.build().as_tuple()
|
|
||||||
# TODO: Move logging out to a custom cursor
|
|
||||||
# logger.debug(
|
|
||||||
# f"Executing query ({query.as_string(cursor)}) with values {values}",
|
|
||||||
# extra={'action': "Query"}
|
|
||||||
# )
|
|
||||||
await cursor.execute(sql.Composed((query,)), values)
|
|
||||||
data = await cursor.fetchall()
|
|
||||||
self.result = self._adapter(*data)
|
|
||||||
return self.result
|
|
||||||
|
|
||||||
async def execute(self, cursor=None) -> QueryResult:
|
|
||||||
"""
|
|
||||||
Execute the query, optionally with the provided cursor, and return the result rows.
|
|
||||||
If no cursor is provided, and no cursor has been set with `with_cursor`,
|
|
||||||
the execution will create a new cursor from the connection and close it automatically.
|
|
||||||
"""
|
|
||||||
# Create a cursor if possible
|
|
||||||
cursor = cursor if cursor is not None else self.cursor
|
|
||||||
if self.cursor is None:
|
|
||||||
if self.conn is None:
|
|
||||||
if self.connector is None:
|
|
||||||
raise ValueError("Cannot execute query without cursor, connection, or connector.")
|
|
||||||
else:
|
|
||||||
async with self.connector.connection() as conn:
|
|
||||||
async with conn.cursor() as cursor:
|
|
||||||
data = await self._execute(cursor)
|
|
||||||
else:
|
|
||||||
async with self.conn.cursor() as cursor:
|
|
||||||
data = await self._execute(cursor)
|
|
||||||
else:
|
|
||||||
data = await self._execute(cursor)
|
|
||||||
return data
|
|
||||||
|
|
||||||
def __await__(self):
|
|
||||||
return self.execute().__await__()
|
|
||||||
|
|
||||||
|
|
||||||
class TableQuery(Query[QueryResult]):
|
|
||||||
"""
|
|
||||||
ABC for an executable query statement expected to be run on a single table.
|
|
||||||
"""
|
|
||||||
__slots__ = (
|
|
||||||
'tableid',
|
|
||||||
'condition', '_extra', '_limit', '_order', '_joins', '_from', '_group'
|
|
||||||
)
|
|
||||||
|
|
||||||
def __init__(self, tableid, *args, **kwargs):
|
|
||||||
super().__init__(*args, **kwargs)
|
|
||||||
self.tableid: sql.Identifier = tableid
|
|
||||||
|
|
||||||
def options(self, **kwargs):
|
|
||||||
"""
|
|
||||||
Set some query options.
|
|
||||||
Default implementation does nothing.
|
|
||||||
Should be overridden to provide specific options.
|
|
||||||
"""
|
|
||||||
return self
|
|
||||||
|
|
||||||
|
|
||||||
class WhereMixin(TableQuery[QueryResult]):
|
|
||||||
__slots__ = ()
|
|
||||||
|
|
||||||
def __init__(self, *args, **kwargs):
|
|
||||||
super().__init__(*args, **kwargs)
|
|
||||||
self.condition: Optional[Condition] = None
|
|
||||||
|
|
||||||
def where(self, *args: Condition, **kwargs):
|
|
||||||
"""
|
|
||||||
Add a Condition to the query.
|
|
||||||
Position arguments should be Conditions,
|
|
||||||
and keyword arguments should be of the form `column=Value`,
|
|
||||||
where Value may be a Value-type or a literal value.
|
|
||||||
All provided Conditions will be and-ed together to create a new Condition.
|
|
||||||
TODO: Maybe just pass this verbatim to a condition.
|
|
||||||
"""
|
|
||||||
if args or kwargs:
|
|
||||||
condition = Condition.construct(*args, **kwargs)
|
|
||||||
if self.condition is not None:
|
|
||||||
condition = self.condition & condition
|
|
||||||
|
|
||||||
self.condition = condition
|
|
||||||
|
|
||||||
return self
|
|
||||||
|
|
||||||
@property
|
|
||||||
def _where_section(self) -> Optional[Expression]:
|
|
||||||
if self.condition is not None:
|
|
||||||
return RawExpr.join_tuples((sql.SQL('WHERE'), ()), self.condition.as_tuple())
|
|
||||||
else:
|
|
||||||
return None
|
|
||||||
|
|
||||||
|
|
||||||
class JOINTYPE(Enum):
|
|
||||||
LEFT = sql.SQL('LEFT JOIN')
|
|
||||||
RIGHT = sql.SQL('RIGHT JOIN')
|
|
||||||
INNER = sql.SQL('INNER JOIN')
|
|
||||||
OUTER = sql.SQL('OUTER JOIN')
|
|
||||||
FULLOUTER = sql.SQL('FULL OUTER JOIN')
|
|
||||||
|
|
||||||
|
|
||||||
class JoinMixin(TableQuery[QueryResult]):
|
|
||||||
__slots__ = ()
|
|
||||||
# TODO: Remember to add join slots to TableQuery
|
|
||||||
|
|
||||||
def __init__(self, *args, **kwargs):
|
|
||||||
super().__init__(*args, **kwargs)
|
|
||||||
self._joins: list[Expression] = []
|
|
||||||
|
|
||||||
def join(self,
|
|
||||||
target: Union[str, Expression],
|
|
||||||
on: Optional[Condition] = None, using: Optional[Union[Expression, tuple[str, ...]]] = None,
|
|
||||||
join_type: JOINTYPE = JOINTYPE.INNER,
|
|
||||||
natural=False):
|
|
||||||
available = (on is not None) + (using is not None) + natural
|
|
||||||
if available == 0:
|
|
||||||
raise ValueError("No conditions given for Query Join")
|
|
||||||
if available > 1:
|
|
||||||
raise ValueError("Exactly one join format must be given for Query Join")
|
|
||||||
|
|
||||||
sections: list[tuple[sql.Composable, tuple[Any, ...]]] = [(join_type.value, ())]
|
|
||||||
if isinstance(target, str):
|
|
||||||
sections.append((sql.Identifier(target), ()))
|
|
||||||
else:
|
|
||||||
sections.append(target.as_tuple())
|
|
||||||
|
|
||||||
if on is not None:
|
|
||||||
sections.append((sql.SQL('ON'), ()))
|
|
||||||
sections.append(on.as_tuple())
|
|
||||||
elif using is not None:
|
|
||||||
sections.append((sql.SQL('USING'), ()))
|
|
||||||
if isinstance(using, Expression):
|
|
||||||
sections.append(using.as_tuple())
|
|
||||||
elif isinstance(using, tuple) and len(using) > 0 and isinstance(using[0], str):
|
|
||||||
cols = sql.SQL("({})").format(sql.SQL(',').join(sql.Identifier(col) for col in using))
|
|
||||||
sections.append((cols, ()))
|
|
||||||
else:
|
|
||||||
raise ValueError("Unrecognised 'using' type.")
|
|
||||||
elif natural:
|
|
||||||
sections.insert(0, (sql.SQL('NATURAL'), ()))
|
|
||||||
|
|
||||||
expr = RawExpr.join_tuples(*sections)
|
|
||||||
self._joins.append(expr)
|
|
||||||
return self
|
|
||||||
|
|
||||||
def leftjoin(self, *args, **kwargs):
|
|
||||||
return self.join(*args, join_type=JOINTYPE.LEFT, **kwargs)
|
|
||||||
|
|
||||||
@property
|
|
||||||
def _join_section(self) -> Optional[Expression]:
|
|
||||||
if self._joins:
|
|
||||||
return RawExpr.join(*self._joins)
|
|
||||||
else:
|
|
||||||
return None
|
|
||||||
|
|
||||||
|
|
||||||
class ExtraMixin(TableQuery[QueryResult]):
|
|
||||||
__slots__ = ()
|
|
||||||
|
|
||||||
def __init__(self, *args, **kwargs):
|
|
||||||
super().__init__(*args, **kwargs)
|
|
||||||
self._extra: Optional[Expression] = None
|
|
||||||
|
|
||||||
def extra(self, extra: sql.Composable, values: tuple[Any, ...] = ()):
|
|
||||||
"""
|
|
||||||
Add an extra string, and optionally values, to this query.
|
|
||||||
The extra string is inserted after any condition, and before the limit.
|
|
||||||
"""
|
|
||||||
extra_expr = RawExpr(extra, values)
|
|
||||||
if self._extra is not None:
|
|
||||||
extra_expr = RawExpr.join(self._extra, extra_expr)
|
|
||||||
self._extra = extra_expr
|
|
||||||
return self
|
|
||||||
|
|
||||||
@property
|
|
||||||
def _extra_section(self) -> Optional[Expression]:
|
|
||||||
if self._extra is None:
|
|
||||||
return None
|
|
||||||
else:
|
|
||||||
return self._extra
|
|
||||||
|
|
||||||
|
|
||||||
class LimitMixin(TableQuery[QueryResult]):
|
|
||||||
__slots__ = ()
|
|
||||||
|
|
||||||
def __init__(self, *args, **kwargs):
|
|
||||||
super().__init__(*args, **kwargs)
|
|
||||||
|
|
||||||
self._limit: Optional[int] = None
|
|
||||||
|
|
||||||
def limit(self, limit: int):
|
|
||||||
"""
|
|
||||||
Add a limit to this query.
|
|
||||||
"""
|
|
||||||
self._limit = limit
|
|
||||||
return self
|
|
||||||
|
|
||||||
@property
|
|
||||||
def _limit_section(self) -> Optional[Expression]:
|
|
||||||
if self._limit is not None:
|
|
||||||
return RawExpr(sql.SQL("LIMIT {}").format(sql.Placeholder()), (self._limit,))
|
|
||||||
else:
|
|
||||||
return None
|
|
||||||
|
|
||||||
|
|
||||||
class FromMixin(TableQuery[QueryResult]):
|
|
||||||
__slots__ = ()
|
|
||||||
|
|
||||||
def __init__(self, *args, **kwargs):
|
|
||||||
super().__init__(*args, **kwargs)
|
|
||||||
self._from: Optional[Expression] = None
|
|
||||||
|
|
||||||
def from_expr(self, _from: Expression):
|
|
||||||
self._from = _from
|
|
||||||
return self
|
|
||||||
|
|
||||||
@property
|
|
||||||
def _from_section(self) -> Optional[Expression]:
|
|
||||||
if self._from is not None:
|
|
||||||
expr, values = self._from.as_tuple()
|
|
||||||
return RawExpr(sql.SQL("FROM {}").format(expr), values)
|
|
||||||
else:
|
|
||||||
return None
|
|
||||||
|
|
||||||
|
|
||||||
class ORDER(Enum):
|
|
||||||
ASC = sql.SQL('ASC')
|
|
||||||
DESC = sql.SQL('DESC')
|
|
||||||
|
|
||||||
|
|
||||||
class NULLS(Enum):
|
|
||||||
FIRST = sql.SQL('NULLS FIRST')
|
|
||||||
LAST = sql.SQL('NULLS LAST')
|
|
||||||
|
|
||||||
|
|
||||||
class OrderMixin(TableQuery[QueryResult]):
|
|
||||||
__slots__ = ()
|
|
||||||
|
|
||||||
def __init__(self, *args, **kwargs):
|
|
||||||
super().__init__(*args, **kwargs)
|
|
||||||
|
|
||||||
self._order: list[Expression] = []
|
|
||||||
|
|
||||||
def order_by(self, expr: Union[Expression, str], direction: Optional[ORDER] = None, nulls: Optional[NULLS] = None):
|
|
||||||
"""
|
|
||||||
Add a single sort expression to the query.
|
|
||||||
This method stacks.
|
|
||||||
"""
|
|
||||||
if isinstance(expr, Expression):
|
|
||||||
string, values = expr.as_tuple()
|
|
||||||
else:
|
|
||||||
string = sql.Identifier(expr)
|
|
||||||
values = ()
|
|
||||||
|
|
||||||
parts = [string]
|
|
||||||
if direction is not None:
|
|
||||||
parts.append(direction.value)
|
|
||||||
if nulls is not None:
|
|
||||||
parts.append(nulls.value)
|
|
||||||
|
|
||||||
order_string = sql.SQL(' ').join(parts)
|
|
||||||
self._order.append(RawExpr(order_string, values))
|
|
||||||
return self
|
|
||||||
|
|
||||||
@property
|
|
||||||
def _order_section(self) -> Optional[Expression]:
|
|
||||||
if self._order:
|
|
||||||
expr = RawExpr.join(*self._order, joiner=sql.SQL(', '))
|
|
||||||
expr.expr = sql.SQL("ORDER BY {}").format(expr.expr)
|
|
||||||
return expr
|
|
||||||
else:
|
|
||||||
return None
|
|
||||||
|
|
||||||
|
|
||||||
class GroupMixin(TableQuery[QueryResult]):
|
|
||||||
__slots__ = ()
|
|
||||||
|
|
||||||
def __init__(self, *args, **kwargs):
|
|
||||||
super().__init__(*args, **kwargs)
|
|
||||||
|
|
||||||
self._group: list[Expression] = []
|
|
||||||
|
|
||||||
def group_by(self, *exprs: Union[Expression, str]):
|
|
||||||
"""
|
|
||||||
Add a group expression(s) to the query.
|
|
||||||
This method stacks.
|
|
||||||
"""
|
|
||||||
for expr in exprs:
|
|
||||||
if isinstance(expr, Expression):
|
|
||||||
self._group.append(expr)
|
|
||||||
else:
|
|
||||||
self._group.append(RawExpr(sql.Identifier(expr)))
|
|
||||||
return self
|
|
||||||
|
|
||||||
@property
|
|
||||||
def _group_section(self) -> Optional[Expression]:
|
|
||||||
if self._group:
|
|
||||||
expr = RawExpr.join(*self._group, joiner=sql.SQL(', '))
|
|
||||||
expr.expr = sql.SQL("GROUP BY {}").format(expr.expr)
|
|
||||||
return expr
|
|
||||||
else:
|
|
||||||
return None
|
|
||||||
|
|
||||||
|
|
||||||
class Insert(ExtraMixin, TableQuery[QueryResult]):
|
|
||||||
"""
|
|
||||||
Query type representing a table insert query.
|
|
||||||
"""
|
|
||||||
# TODO: Support ON CONFLICT for upserts
|
|
||||||
__slots__ = ('_columns', '_values', '_conflict')
|
|
||||||
|
|
||||||
def __init__(self, *args, **kwargs):
|
|
||||||
super().__init__(*args, **kwargs)
|
|
||||||
self._columns: tuple[str, ...] = ()
|
|
||||||
self._values: tuple[tuple[Any, ...], ...] = ()
|
|
||||||
self._conflict: Optional[Expression] = None
|
|
||||||
|
|
||||||
def insert(self, columns, *values):
|
|
||||||
"""
|
|
||||||
Insert the given data.
|
|
||||||
|
|
||||||
Parameters
|
|
||||||
----------
|
|
||||||
columns: tuple[str]
|
|
||||||
Tuple of column names to insert.
|
|
||||||
|
|
||||||
values: tuple[tuple[Any, ...], ...]
|
|
||||||
Tuple of values to insert, corresponding to the columns.
|
|
||||||
"""
|
|
||||||
if not values:
|
|
||||||
raise ValueError("Cannot insert zero rows.")
|
|
||||||
if len(values[0]) != len(columns):
|
|
||||||
raise ValueError("Number of columns does not match length of values.")
|
|
||||||
|
|
||||||
self._columns = columns
|
|
||||||
self._values = values
|
|
||||||
return self
|
|
||||||
|
|
||||||
def on_conflict(self, ignore=False):
|
|
||||||
# TODO lots more we can do here
|
|
||||||
# Maybe return a Conflict object that can chain itself (not the query)
|
|
||||||
if ignore:
|
|
||||||
self._conflict = RawExpr(sql.SQL('DO NOTHING'))
|
|
||||||
return self
|
|
||||||
|
|
||||||
@property
|
|
||||||
def _conflict_section(self) -> Optional[Expression]:
|
|
||||||
if self._conflict is not None:
|
|
||||||
e, v = self._conflict.as_tuple()
|
|
||||||
expr = RawExpr(
|
|
||||||
sql.SQL("ON CONFLICT {}").format(
|
|
||||||
e
|
|
||||||
),
|
|
||||||
v
|
|
||||||
)
|
|
||||||
return expr
|
|
||||||
return None
|
|
||||||
|
|
||||||
def build(self):
|
|
||||||
columns = sql.SQL(',').join(map(sql.Identifier, self._columns))
|
|
||||||
single_value_str = sql.SQL('({})').format(
|
|
||||||
sql.SQL(',').join(sql.Placeholder() * len(self._columns))
|
|
||||||
)
|
|
||||||
values_str = sql.SQL(',').join(single_value_str * len(self._values))
|
|
||||||
|
|
||||||
# TODO: Check efficiency of inserting multiple values like this
|
|
||||||
# Also implement a Copy query
|
|
||||||
base = sql.SQL("INSERT INTO {table} ({columns}) VALUES {values_str}").format(
|
|
||||||
table=self.tableid,
|
|
||||||
columns=columns,
|
|
||||||
values_str=values_str
|
|
||||||
)
|
|
||||||
|
|
||||||
sections = [
|
|
||||||
RawExpr(base, tuple(chain(*self._values))),
|
|
||||||
self._conflict_section,
|
|
||||||
self._extra_section,
|
|
||||||
RawExpr(sql.SQL('RETURNING *'))
|
|
||||||
]
|
|
||||||
|
|
||||||
sections = (section for section in sections if section is not None)
|
|
||||||
return RawExpr.join(*sections)
|
|
||||||
|
|
||||||
|
|
||||||
class Select(WhereMixin, ExtraMixin, OrderMixin, LimitMixin, JoinMixin, GroupMixin, TableQuery[QueryResult]):
|
|
||||||
"""
|
|
||||||
Select rows from a table matching provided conditions.
|
|
||||||
"""
|
|
||||||
__slots__ = ('_columns',)
|
|
||||||
|
|
||||||
def __init__(self, *args, **kwargs):
|
|
||||||
super().__init__(*args, **kwargs)
|
|
||||||
self._columns: tuple[Expression, ...] = ()
|
|
||||||
|
|
||||||
def select(self, *columns: str, **exprs: Union[str, sql.Composable, Expression]):
|
|
||||||
"""
|
|
||||||
Set the columns and expressions to select.
|
|
||||||
If none are given, selects all columns.
|
|
||||||
"""
|
|
||||||
cols: List[Expression] = []
|
|
||||||
if columns:
|
|
||||||
cols.extend(map(RawExpr, map(sql.Identifier, columns)))
|
|
||||||
if exprs:
|
|
||||||
for name, expr in exprs.items():
|
|
||||||
if isinstance(expr, str):
|
|
||||||
cols.append(
|
|
||||||
RawExpr(sql.SQL(expr) + sql.SQL(' AS ') + sql.Identifier(name))
|
|
||||||
)
|
|
||||||
elif isinstance(expr, sql.Composable):
|
|
||||||
cols.append(
|
|
||||||
RawExpr(expr + sql.SQL(' AS ') + sql.Identifier(name))
|
|
||||||
)
|
|
||||||
elif isinstance(expr, Expression):
|
|
||||||
value_expr, value_values = expr.as_tuple()
|
|
||||||
cols.append(RawExpr(
|
|
||||||
value_expr + sql.SQL(' AS ') + sql.Identifier(name),
|
|
||||||
value_values
|
|
||||||
))
|
|
||||||
if cols:
|
|
||||||
self._columns = (*self._columns, *cols)
|
|
||||||
return self
|
|
||||||
|
|
||||||
def build(self):
|
|
||||||
if not self._columns:
|
|
||||||
columns, columns_values = sql.SQL('*'), ()
|
|
||||||
else:
|
|
||||||
columns, columns_values = RawExpr.join(*self._columns, joiner=sql.SQL(',')).as_tuple()
|
|
||||||
|
|
||||||
base = sql.SQL("SELECT {columns} FROM {table}").format(
|
|
||||||
columns=columns,
|
|
||||||
table=self.tableid
|
|
||||||
)
|
|
||||||
|
|
||||||
sections = [
|
|
||||||
RawExpr(base, columns_values),
|
|
||||||
self._join_section,
|
|
||||||
self._where_section,
|
|
||||||
self._group_section,
|
|
||||||
self._extra_section,
|
|
||||||
self._order_section,
|
|
||||||
self._limit_section,
|
|
||||||
]
|
|
||||||
|
|
||||||
sections = (section for section in sections if section is not None)
|
|
||||||
return RawExpr.join(*sections)
|
|
||||||
|
|
||||||
|
|
||||||
class Delete(WhereMixin, ExtraMixin, TableQuery[QueryResult]):
|
|
||||||
"""
|
|
||||||
Query type representing a table delete query.
|
|
||||||
"""
|
|
||||||
# TODO: Cascade option for delete, maybe other options
|
|
||||||
# TODO: Require a where unless specifically disabled, for safety
|
|
||||||
|
|
||||||
def build(self):
|
|
||||||
base = sql.SQL("DELETE FROM {table}").format(
|
|
||||||
table=self.tableid,
|
|
||||||
)
|
|
||||||
sections = [
|
|
||||||
RawExpr(base),
|
|
||||||
self._where_section,
|
|
||||||
self._extra_section,
|
|
||||||
RawExpr(sql.SQL('RETURNING *'))
|
|
||||||
]
|
|
||||||
|
|
||||||
sections = (section for section in sections if section is not None)
|
|
||||||
return RawExpr.join(*sections)
|
|
||||||
|
|
||||||
|
|
||||||
class Update(LimitMixin, WhereMixin, ExtraMixin, FromMixin, TableQuery[QueryResult]):
|
|
||||||
__slots__ = (
|
|
||||||
'_set',
|
|
||||||
)
|
|
||||||
# TODO: Again, require a where unless specifically disabled
|
|
||||||
|
|
||||||
def __init__(self, *args, **kwargs):
|
|
||||||
super().__init__(*args, **kwargs)
|
|
||||||
self._set: List[Expression] = []
|
|
||||||
|
|
||||||
def set(self, **column_values: Union[Any, Expression]):
|
|
||||||
exprs: List[Expression] = []
|
|
||||||
for name, value in column_values.items():
|
|
||||||
if isinstance(value, Expression):
|
|
||||||
value_tup = value.as_tuple()
|
|
||||||
else:
|
|
||||||
value_tup = (sql.Placeholder(), (value,))
|
|
||||||
|
|
||||||
exprs.append(
|
|
||||||
RawExpr.join_tuples(
|
|
||||||
(sql.Identifier(name), ()),
|
|
||||||
value_tup,
|
|
||||||
joiner=sql.SQL(' = ')
|
|
||||||
)
|
|
||||||
)
|
|
||||||
self._set.extend(exprs)
|
|
||||||
return self
|
|
||||||
|
|
||||||
def build(self):
|
|
||||||
if not self._set:
|
|
||||||
raise ValueError("No columns provided to update.")
|
|
||||||
set_expr, set_values = RawExpr.join(*self._set, joiner=sql.SQL(', ')).as_tuple()
|
|
||||||
|
|
||||||
base = sql.SQL("UPDATE {table} SET {set}").format(
|
|
||||||
table=self.tableid,
|
|
||||||
set=set_expr
|
|
||||||
)
|
|
||||||
sections = [
|
|
||||||
RawExpr(base, set_values),
|
|
||||||
self._from_section,
|
|
||||||
self._where_section,
|
|
||||||
self._extra_section,
|
|
||||||
self._limit_section,
|
|
||||||
RawExpr(sql.SQL('RETURNING *'))
|
|
||||||
]
|
|
||||||
|
|
||||||
sections = (section for section in sections if section is not None)
|
|
||||||
return RawExpr.join(*sections)
|
|
||||||
|
|
||||||
|
|
||||||
# async def upsert(cursor, table, constraint, **values):
|
|
||||||
# """
|
|
||||||
# Insert or on conflict update.
|
|
||||||
# """
|
|
||||||
# valuedict = values
|
|
||||||
# keys, values = zip(*values.items())
|
|
||||||
#
|
|
||||||
# key_str = _format_insertkeys(keys)
|
|
||||||
# value_str, values = _format_insertvalues(values)
|
|
||||||
# update_key_str, update_key_values = _format_updatestr(valuedict)
|
|
||||||
#
|
|
||||||
# if not isinstance(constraint, str):
|
|
||||||
# constraint = ", ".join(constraint)
|
|
||||||
#
|
|
||||||
# await cursor.execute(
|
|
||||||
# 'INSERT INTO {} {} VALUES {} ON CONFLICT({}) DO UPDATE SET {} RETURNING *'.format(
|
|
||||||
# table, key_str, value_str, constraint, update_key_str
|
|
||||||
# ),
|
|
||||||
# tuple((*values, *update_key_values))
|
|
||||||
# )
|
|
||||||
# return await cursor.fetchone()
|
|
||||||
|
|
||||||
|
|
||||||
# def update_many(table, *values, set_keys=None, where_keys=None, cast_row=None, cursor=None):
|
|
||||||
# cursor = cursor or conn.cursor()
|
|
||||||
#
|
|
||||||
# # TODO: executemany or copy syntax now
|
|
||||||
# return execute_values(
|
|
||||||
# cursor,
|
|
||||||
# """
|
|
||||||
# UPDATE {table}
|
|
||||||
# SET {set_clause}
|
|
||||||
# FROM (VALUES {cast_row}%s)
|
|
||||||
# AS {temp_table}
|
|
||||||
# WHERE {where_clause}
|
|
||||||
# RETURNING *
|
|
||||||
# """.format(
|
|
||||||
# table=table,
|
|
||||||
# set_clause=', '.join("{0} = _t.{0}".format(key) for key in set_keys),
|
|
||||||
# cast_row=cast_row + ',' if cast_row else '',
|
|
||||||
# where_clause=' AND '.join("{1}.{0} = _t.{0}".format(key, table) for key in where_keys),
|
|
||||||
# temp_table="_t ({})".format(', '.join(set_keys + where_keys))
|
|
||||||
# ),
|
|
||||||
# values,
|
|
||||||
# fetch=True
|
|
||||||
# )
|
|
||||||
@@ -1,102 +0,0 @@
|
|||||||
from typing import Protocol, runtime_checkable, Optional
|
|
||||||
|
|
||||||
from psycopg import AsyncConnection
|
|
||||||
|
|
||||||
from .connector import Connector, Connectable
|
|
||||||
|
|
||||||
|
|
||||||
@runtime_checkable
|
|
||||||
class _Attachable(Connectable, Protocol):
|
|
||||||
def attach_to(self, registry: 'Registry'):
|
|
||||||
raise NotImplementedError
|
|
||||||
|
|
||||||
|
|
||||||
class Registry:
|
|
||||||
_attached: list[_Attachable] = []
|
|
||||||
_name: Optional[str] = None
|
|
||||||
|
|
||||||
def __init_subclass__(cls, name=None):
|
|
||||||
attached = []
|
|
||||||
for _, member in cls.__dict__.items():
|
|
||||||
if isinstance(member, _Attachable):
|
|
||||||
attached.append(member)
|
|
||||||
cls._attached = attached
|
|
||||||
cls._name = name or cls.__name__
|
|
||||||
|
|
||||||
def __init__(self, name=None):
|
|
||||||
self._conn: Optional[Connector] = None
|
|
||||||
self.name: str = name if name is not None else self._name
|
|
||||||
if self.name is None:
|
|
||||||
raise ValueError("A Registry must have a name!")
|
|
||||||
|
|
||||||
self.init_tasks = []
|
|
||||||
|
|
||||||
for member in self._attached:
|
|
||||||
member.attach_to(self)
|
|
||||||
|
|
||||||
def bind(self, connector: Connector):
|
|
||||||
self._conn = connector
|
|
||||||
for child in self._attached:
|
|
||||||
child.bind(connector)
|
|
||||||
|
|
||||||
def attach(self, attachable):
|
|
||||||
self._attached.append(attachable)
|
|
||||||
if self._conn is not None:
|
|
||||||
attachable.bind(self._conn)
|
|
||||||
return attachable
|
|
||||||
|
|
||||||
def init_task(self, coro):
|
|
||||||
"""
|
|
||||||
Initialisation tasks are run to setup the registry state.
|
|
||||||
These tasks will be run in the event loop, after connection to the database.
|
|
||||||
These tasks should be idempotent, as they may be run on reload and reconnect.
|
|
||||||
"""
|
|
||||||
self.init_tasks.append(coro)
|
|
||||||
return coro
|
|
||||||
|
|
||||||
async def init(self):
|
|
||||||
for task in self.init_tasks:
|
|
||||||
await task(self)
|
|
||||||
return self
|
|
||||||
|
|
||||||
|
|
||||||
class AttachableClass:
|
|
||||||
"""ABC for a default implementation of an Attachable class."""
|
|
||||||
|
|
||||||
_connector: Optional[Connector] = None
|
|
||||||
_registry: Optional[Registry] = None
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def bind(cls, connector: Connector):
|
|
||||||
cls._connector = connector
|
|
||||||
connector.connect_hook(cls.on_connect)
|
|
||||||
return cls
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def attach_to(cls, registry: Registry):
|
|
||||||
cls._registry = registry
|
|
||||||
return cls
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
async def on_connect(cls, connection: AsyncConnection):
|
|
||||||
pass
|
|
||||||
|
|
||||||
|
|
||||||
class Attachable:
|
|
||||||
"""ABC for a default implementation of an Attachable object."""
|
|
||||||
|
|
||||||
def __init__(self, *args, **kwargs):
|
|
||||||
self._connector: Optional[Connector] = None
|
|
||||||
self._registry: Optional[Registry] = None
|
|
||||||
|
|
||||||
def bind(self, connector: Connector):
|
|
||||||
self._connector = connector
|
|
||||||
connector.connect_hook(self.on_connect)
|
|
||||||
return self
|
|
||||||
|
|
||||||
def attach_to(self, registry: Registry):
|
|
||||||
self._registry = registry
|
|
||||||
return self
|
|
||||||
|
|
||||||
async def on_connect(self, connection: AsyncConnection):
|
|
||||||
pass
|
|
||||||
@@ -1,95 +0,0 @@
|
|||||||
from typing import Optional
|
|
||||||
from psycopg.rows import DictRow
|
|
||||||
from psycopg import sql
|
|
||||||
|
|
||||||
from . import queries as q
|
|
||||||
from .connector import Connector
|
|
||||||
from .registry import Registry
|
|
||||||
|
|
||||||
|
|
||||||
class Table:
|
|
||||||
"""
|
|
||||||
Transparent interface to a single table structure in the database.
|
|
||||||
Contains standard methods to access the table.
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(self, name, *args, schema='public', **kwargs):
|
|
||||||
self.name: str = name
|
|
||||||
self.schema: str = schema
|
|
||||||
self.connector: Connector = None
|
|
||||||
|
|
||||||
@property
|
|
||||||
def identifier(self):
|
|
||||||
if self.schema == 'public':
|
|
||||||
return sql.Identifier(self.name)
|
|
||||||
else:
|
|
||||||
return sql.Identifier(self.schema, self.name)
|
|
||||||
|
|
||||||
def bind(self, connector: Connector):
|
|
||||||
self.connector = connector
|
|
||||||
return self
|
|
||||||
|
|
||||||
def attach_to(self, registry: Registry):
|
|
||||||
self._registry = registry
|
|
||||||
return self
|
|
||||||
|
|
||||||
def _many_query_adapter(self, *data: DictRow) -> tuple[DictRow, ...]:
|
|
||||||
return data
|
|
||||||
|
|
||||||
def _single_query_adapter(self, *data: DictRow) -> Optional[DictRow]:
|
|
||||||
if data:
|
|
||||||
return data[0]
|
|
||||||
else:
|
|
||||||
return None
|
|
||||||
|
|
||||||
def _delete_query_adapter(self, *data: DictRow) -> tuple[DictRow, ...]:
|
|
||||||
return data
|
|
||||||
|
|
||||||
def select_where(self, *args, **kwargs) -> q.Select[tuple[DictRow, ...]]:
|
|
||||||
return q.Select(
|
|
||||||
self.identifier,
|
|
||||||
row_adapter=self._many_query_adapter,
|
|
||||||
connector=self.connector
|
|
||||||
).where(*args, **kwargs)
|
|
||||||
|
|
||||||
def select_one_where(self, *args, **kwargs) -> q.Select[DictRow]:
|
|
||||||
return q.Select(
|
|
||||||
self.identifier,
|
|
||||||
row_adapter=self._single_query_adapter,
|
|
||||||
connector=self.connector
|
|
||||||
).where(*args, **kwargs)
|
|
||||||
|
|
||||||
def update_where(self, *args, **kwargs) -> q.Update[tuple[DictRow, ...]]:
|
|
||||||
return q.Update(
|
|
||||||
self.identifier,
|
|
||||||
row_adapter=self._many_query_adapter,
|
|
||||||
connector=self.connector
|
|
||||||
).where(*args, **kwargs)
|
|
||||||
|
|
||||||
def delete_where(self, *args, **kwargs) -> q.Delete[tuple[DictRow, ...]]:
|
|
||||||
return q.Delete(
|
|
||||||
self.identifier,
|
|
||||||
row_adapter=self._many_query_adapter,
|
|
||||||
connector=self.connector
|
|
||||||
).where(*args, **kwargs)
|
|
||||||
|
|
||||||
def insert(self, **column_values) -> q.Insert[DictRow]:
|
|
||||||
return q.Insert(
|
|
||||||
self.identifier,
|
|
||||||
row_adapter=self._single_query_adapter,
|
|
||||||
connector=self.connector
|
|
||||||
).insert(column_values.keys(), column_values.values())
|
|
||||||
|
|
||||||
def insert_many(self, *args, **kwargs) -> q.Insert[tuple[DictRow, ...]]:
|
|
||||||
return q.Insert(
|
|
||||||
self.identifier,
|
|
||||||
row_adapter=self._many_query_adapter,
|
|
||||||
connector=self.connector
|
|
||||||
).insert(*args, **kwargs)
|
|
||||||
|
|
||||||
# def update_many(self, *args, **kwargs):
|
|
||||||
# with self.conn:
|
|
||||||
# return update_many(self.identifier, *args, **kwargs)
|
|
||||||
|
|
||||||
# def upsert(self, *args, **kwargs):
|
|
||||||
# return upsert(self.identifier, *args, **kwargs)
|
|
||||||
+42
-2
@@ -3,6 +3,7 @@ import logging
|
|||||||
import asyncio
|
import asyncio
|
||||||
from weakref import WeakValueDictionary
|
from weakref import WeakValueDictionary
|
||||||
|
|
||||||
|
from constants import SCHEMA_VERSIONS
|
||||||
import discord
|
import discord
|
||||||
from discord.utils import MISSING
|
from discord.utils import MISSING
|
||||||
from discord.ext.commands import Bot, Cog, HybridCommand, HybridCommandError
|
from discord.ext.commands import Bot, Cog, HybridCommand, HybridCommandError
|
||||||
@@ -10,8 +11,10 @@ from discord.ext.commands.errors import CommandInvokeError, CheckFailure
|
|||||||
from discord.app_commands.errors import CommandInvokeError as appCommandInvokeError, TransformerError
|
from discord.app_commands.errors import CommandInvokeError as appCommandInvokeError, TransformerError
|
||||||
from aiohttp import ClientSession
|
from aiohttp import ClientSession
|
||||||
|
|
||||||
from data import Database
|
from data import Database, ORDER
|
||||||
from utils.lib import tabulate
|
from utils.lib import tabulate
|
||||||
|
from babel.translator import LeoBabel
|
||||||
|
from botdata import BotData, VersionHistory
|
||||||
|
|
||||||
from .config import Conf
|
from .config import Conf
|
||||||
from .logger import logging_context, log_context, log_action_stack, log_wrap, set_logging_context
|
from .logger import logging_context, log_context, log_action_stack, log_wrap, set_logging_context
|
||||||
@@ -23,6 +26,7 @@ from .monitor import SystemMonitor, ComponentMonitor, StatusLevel, ComponentStat
|
|||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from core.cog import CoreCog
|
from core.cog import CoreCog
|
||||||
|
from modules.profiles.profiles.discord.cog import ProfilesCog
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -42,7 +46,9 @@ class LionBot(Bot):
|
|||||||
self.appname = appname
|
self.appname = appname
|
||||||
self.shardname = shardname
|
self.shardname = shardname
|
||||||
# self.appdata = appdata
|
# self.appdata = appdata
|
||||||
|
self.data: BotData = db.load_registry(BotData())
|
||||||
self.config = config
|
self.config = config
|
||||||
|
self.translator = LeoBabel()
|
||||||
|
|
||||||
self.system_monitor = SystemMonitor()
|
self.system_monitor = SystemMonitor()
|
||||||
self.monitor = ComponentMonitor('LionBot', self._monitor_status)
|
self.monitor = ComponentMonitor('LionBot', self._monitor_status)
|
||||||
@@ -51,10 +57,18 @@ class LionBot(Bot):
|
|||||||
self._locks = WeakValueDictionary()
|
self._locks = WeakValueDictionary()
|
||||||
self._running_events = set()
|
self._running_events = set()
|
||||||
|
|
||||||
|
@property
|
||||||
|
def dbconn(self):
|
||||||
|
return self.db
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def core(self):
|
def core(self):
|
||||||
return self.get_cog('CoreCog')
|
return self.get_cog('CoreCog')
|
||||||
|
|
||||||
|
@property
|
||||||
|
def profiles(self):
|
||||||
|
return self.get_cog('ProfilesCog')
|
||||||
|
|
||||||
async def _monitor_status(self):
|
async def _monitor_status(self):
|
||||||
if self.is_closed():
|
if self.is_closed():
|
||||||
level = StatusLevel.ERRORED
|
level = StatusLevel.ERRORED
|
||||||
@@ -101,6 +115,10 @@ class LionBot(Bot):
|
|||||||
def get_cog(self, name: Literal['CoreCog']) -> 'CoreCog':
|
def get_cog(self, name: Literal['CoreCog']) -> 'CoreCog':
|
||||||
...
|
...
|
||||||
|
|
||||||
|
@overload
|
||||||
|
def get_cog(self, name: Literal['ProfilesCog']) -> 'ProfilesCog':
|
||||||
|
...
|
||||||
|
|
||||||
@overload
|
@overload
|
||||||
def get_cog(self, name: str) -> Optional[Cog]:
|
def get_cog(self, name: str) -> Optional[Cog]:
|
||||||
...
|
...
|
||||||
@@ -127,6 +145,10 @@ class LionBot(Bot):
|
|||||||
await wrapper()
|
await wrapper()
|
||||||
|
|
||||||
async def start(self, token: str, *, reconnect: bool = True):
|
async def start(self, token: str, *, reconnect: bool = True):
|
||||||
|
await self.data.init()
|
||||||
|
for component, req in SCHEMA_VERSIONS.items():
|
||||||
|
await self.version_check(component, req)
|
||||||
|
|
||||||
with logging_context(action="Login"):
|
with logging_context(action="Login"):
|
||||||
start_task = asyncio.create_task(self.login(token))
|
start_task = asyncio.create_task(self.login(token))
|
||||||
await start_task
|
await start_task
|
||||||
@@ -135,6 +157,24 @@ class LionBot(Bot):
|
|||||||
run_task = asyncio.create_task(self.connect(reconnect=reconnect))
|
run_task = asyncio.create_task(self.connect(reconnect=reconnect))
|
||||||
await run_task
|
await run_task
|
||||||
|
|
||||||
|
async def version_check(self, component: str, req_version: int):
|
||||||
|
# Query the database to confirm that the given component is listed with the given version.
|
||||||
|
# Typically done upon loading a component
|
||||||
|
rows = await VersionHistory.fetch_where(component=component).order_by('_timestamp', ORDER.DESC).limit(1)
|
||||||
|
|
||||||
|
version = rows[0].to_version if rows else 0
|
||||||
|
|
||||||
|
if version != req_version:
|
||||||
|
raise ValueError(f"Component {component} failed version check. Has version '{version}', required version '{req_version}'")
|
||||||
|
else:
|
||||||
|
logger.debug(
|
||||||
|
"Component %s passed version check with version %s",
|
||||||
|
component,
|
||||||
|
version
|
||||||
|
)
|
||||||
|
return True
|
||||||
|
|
||||||
|
|
||||||
def dispatch(self, event_name: str, *args, **kwargs):
|
def dispatch(self, event_name: str, *args, **kwargs):
|
||||||
with logging_context(action=f"Dispatch {event_name}"):
|
with logging_context(action=f"Dispatch {event_name}"):
|
||||||
super().dispatch(event_name, *args, **kwargs)
|
super().dispatch(event_name, *args, **kwargs)
|
||||||
@@ -189,7 +229,7 @@ class LionBot(Bot):
|
|||||||
# TODO: Some of these could have more user-feedback
|
# TODO: Some of these could have more user-feedback
|
||||||
logger.debug(f"Handling command error for {ctx}: {exception}")
|
logger.debug(f"Handling command error for {ctx}: {exception}")
|
||||||
if isinstance(ctx.command, HybridCommand) and ctx.command.app_command:
|
if isinstance(ctx.command, HybridCommand) and ctx.command.app_command:
|
||||||
cmd_str = ctx.command.app_command.to_dict()
|
cmd_str = ctx.command.app_command.to_dict(self.tree)
|
||||||
else:
|
else:
|
||||||
cmd_str = str(ctx.command)
|
cmd_str = str(ctx.command)
|
||||||
try:
|
try:
|
||||||
|
|||||||
@@ -9,6 +9,7 @@ from discord.ext.commands import Context
|
|||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from .LionBot import LionBot
|
from .LionBot import LionBot
|
||||||
|
from modules.profiles.profiles.data import UserProfile, Community
|
||||||
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -46,6 +47,8 @@ class LionContext(Context['LionBot']):
|
|||||||
Extends Context to add Lion-specific methods and attributes.
|
Extends Context to add Lion-specific methods and attributes.
|
||||||
Also adds several contextual wrapped utilities for simpler user during command invocation.
|
Also adds several contextual wrapped utilities for simpler user during command invocation.
|
||||||
"""
|
"""
|
||||||
|
profile: 'UserProfile'
|
||||||
|
community: 'Community'
|
||||||
|
|
||||||
def __repr__(self):
|
def __repr__(self):
|
||||||
parts = {}
|
parts = {}
|
||||||
|
|||||||
@@ -131,7 +131,7 @@ class LionTree(CommandTree):
|
|||||||
return
|
return
|
||||||
|
|
||||||
set_logging_context(action=f"Run {command.qualified_name}")
|
set_logging_context(action=f"Run {command.qualified_name}")
|
||||||
logger.debug(f"Running command '{command.qualified_name}': {command.to_dict()}")
|
logger.debug(f"Running command '{command.qualified_name}': {command.to_dict(self)}")
|
||||||
try:
|
try:
|
||||||
await command._invoke_with_namespace(interaction, namespace)
|
await command._invoke_with_namespace(interaction, namespace)
|
||||||
except AppCommandError as e:
|
except AppCommandError as e:
|
||||||
|
|||||||
@@ -1,8 +1,13 @@
|
|||||||
this_package = 'modules'
|
this_package = "modules"
|
||||||
|
|
||||||
active = [
|
active = [
|
||||||
'.sysadmin',
|
".sysadmin",
|
||||||
'.voicefix',
|
".profiles",
|
||||||
|
".voicefix",
|
||||||
|
".messagelogger",
|
||||||
|
".voicelog",
|
||||||
|
".yarn",
|
||||||
|
".pluscampaign",
|
||||||
]
|
]
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Submodule
+1
Submodule src/modules/messagelogger added at 166e310f96
Submodule
+1
Submodule src/modules/pluscampaign added at a18885511a
Submodule
+1
Submodule src/modules/profiles added at 8818263d88
Submodule
+1
Submodule src/modules/voicefix added at 70d089f5de
@@ -1,449 +0,0 @@
|
|||||||
from collections import defaultdict
|
|
||||||
from typing import Optional
|
|
||||||
import asyncio
|
|
||||||
from cachetools import FIFOCache
|
|
||||||
|
|
||||||
import discord
|
|
||||||
from discord.abc import GuildChannel
|
|
||||||
from discord.ext import commands as cmds
|
|
||||||
from discord import app_commands as appcmds
|
|
||||||
|
|
||||||
from meta import LionBot, LionCog, LionContext
|
|
||||||
from meta.errors import ResponseTimedOut, SafeCancellation, UserInputError
|
|
||||||
from utils.ui import Confirm
|
|
||||||
|
|
||||||
from . import logger
|
|
||||||
from .data import LinkData
|
|
||||||
|
|
||||||
|
|
||||||
async def prepare_attachments(attachments: list[discord.Attachment]):
|
|
||||||
results = []
|
|
||||||
for attach in attachments:
|
|
||||||
try:
|
|
||||||
as_file = await attach.to_file(spoiler=attach.is_spoiler())
|
|
||||||
results.append(as_file)
|
|
||||||
except discord.HTTPException:
|
|
||||||
pass
|
|
||||||
return results
|
|
||||||
|
|
||||||
|
|
||||||
async def prepare_embeds(message: discord.Message):
|
|
||||||
embeds = [embed for embed in message.embeds if embed.type == 'rich']
|
|
||||||
if message.reference:
|
|
||||||
embed = discord.Embed(
|
|
||||||
colour=discord.Colour.dark_gray(),
|
|
||||||
description=f"Reply to {message.reference.jump_url}"
|
|
||||||
)
|
|
||||||
embeds.append(embed)
|
|
||||||
return embeds
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
class VoiceFixCog(LionCog):
|
|
||||||
def __init__(self, bot: LionBot):
|
|
||||||
self.bot = bot
|
|
||||||
self.data = bot.db.load_registry(LinkData())
|
|
||||||
|
|
||||||
# Map of linkids to list of channelids
|
|
||||||
self.link_channels = {}
|
|
||||||
|
|
||||||
# Map of channelids to linkids
|
|
||||||
self.channel_links = {}
|
|
||||||
|
|
||||||
# Map of channelids to initialised discord.Webhook
|
|
||||||
self.hooks = {}
|
|
||||||
|
|
||||||
# Map of messageid to list of (channelid, webhookmsg) pairs, for updates
|
|
||||||
self.message_cache = FIFOCache(maxsize=200)
|
|
||||||
# webhook msgid -> orig msgid
|
|
||||||
self.wmessages = FIFOCache(maxsize=600)
|
|
||||||
|
|
||||||
self.lock = asyncio.Lock()
|
|
||||||
|
|
||||||
|
|
||||||
async def cog_load(self):
|
|
||||||
await self.data.init()
|
|
||||||
|
|
||||||
await self.reload_links()
|
|
||||||
|
|
||||||
async def reload_links(self):
|
|
||||||
records = await self.data.channel_links.select_where()
|
|
||||||
channel_links = defaultdict(set)
|
|
||||||
link_channels = defaultdict(set)
|
|
||||||
|
|
||||||
for record in records:
|
|
||||||
linkid = record['linkid']
|
|
||||||
channelid = record['channelid']
|
|
||||||
|
|
||||||
channel_links[channelid].add(linkid)
|
|
||||||
link_channels[linkid].add(channelid)
|
|
||||||
|
|
||||||
channelids = list(channel_links.keys())
|
|
||||||
if channelids:
|
|
||||||
await self.data.LinkHook.fetch_where(channelid=channelids)
|
|
||||||
for channelid in channelids:
|
|
||||||
# Will hit cache, so don't need any more data queries
|
|
||||||
await self.fetch_webhook_for(channelid)
|
|
||||||
|
|
||||||
self.channel_links = {cid: tuple(linkids) for cid, linkids in channel_links.items()}
|
|
||||||
self.link_channels = {lid: tuple(cids) for lid, cids in link_channels.items()}
|
|
||||||
|
|
||||||
logger.info(
|
|
||||||
f"Loaded '{len(link_channels)}' channel links with '{len(self.channel_links)}' linked channels."
|
|
||||||
)
|
|
||||||
|
|
||||||
@LionCog.listener('on_message')
|
|
||||||
async def on_message(self, message: discord.Message):
|
|
||||||
# Don't need this because everything except explicit messages are webhooks now
|
|
||||||
# if self.bot.user and (message.author.id == self.bot.user.id):
|
|
||||||
# return
|
|
||||||
if message.webhook_id:
|
|
||||||
return
|
|
||||||
|
|
||||||
async with self.lock:
|
|
||||||
sent = []
|
|
||||||
linkids = self.channel_links.get(message.channel.id, ())
|
|
||||||
if linkids:
|
|
||||||
for linkid in linkids:
|
|
||||||
for channelid in self.link_channels[linkid]:
|
|
||||||
if channelid != message.channel.id:
|
|
||||||
if message.attachments:
|
|
||||||
files = await prepare_attachments(message.attachments)
|
|
||||||
else:
|
|
||||||
files = []
|
|
||||||
|
|
||||||
hook = self.hooks[channelid]
|
|
||||||
avatar = message.author.avatar or message.author.default_avatar
|
|
||||||
msg = await hook.send(
|
|
||||||
content=message.content,
|
|
||||||
wait=True,
|
|
||||||
username=message.author.display_name,
|
|
||||||
avatar_url=avatar.url,
|
|
||||||
embeds=await prepare_embeds(message),
|
|
||||||
files=files,
|
|
||||||
allowed_mentions=discord.AllowedMentions.none()
|
|
||||||
)
|
|
||||||
sent.append((channelid, msg))
|
|
||||||
self.wmessages[msg.id] = message.id
|
|
||||||
if sent:
|
|
||||||
# For easier lookup
|
|
||||||
self.wmessages[message.id] = message.id
|
|
||||||
sent.append((message.channel.id, message))
|
|
||||||
|
|
||||||
self.message_cache[message.id] = sent
|
|
||||||
logger.info(f"Forwarded message {message.id}")
|
|
||||||
|
|
||||||
|
|
||||||
@LionCog.listener('on_message_edit')
|
|
||||||
async def on_message_edit(self, before, after):
|
|
||||||
async with self.lock:
|
|
||||||
cached_sent = self.message_cache.pop(before.id, ())
|
|
||||||
new_sent = []
|
|
||||||
for cid, msg in cached_sent:
|
|
||||||
try:
|
|
||||||
if msg.id != before.id:
|
|
||||||
msg = await msg.edit(
|
|
||||||
content=after.content,
|
|
||||||
embeds=await prepare_embeds(after),
|
|
||||||
)
|
|
||||||
new_sent.append((cid, msg))
|
|
||||||
except discord.NotFound:
|
|
||||||
pass
|
|
||||||
if new_sent:
|
|
||||||
self.message_cache[after.id] = new_sent
|
|
||||||
|
|
||||||
@LionCog.listener('on_message_delete')
|
|
||||||
async def on_message_delete(self, message):
|
|
||||||
async with self.lock:
|
|
||||||
origid = self.wmessages.get(message.id, None)
|
|
||||||
if origid:
|
|
||||||
cached_sent = self.message_cache.pop(origid, ())
|
|
||||||
for _, msg in cached_sent:
|
|
||||||
try:
|
|
||||||
if msg.id != message.id:
|
|
||||||
await msg.delete()
|
|
||||||
except discord.NotFound:
|
|
||||||
pass
|
|
||||||
|
|
||||||
@LionCog.listener('on_reaction_add')
|
|
||||||
async def on_reaction_add(self, reaction: discord.Reaction, user: discord.User):
|
|
||||||
async with self.lock:
|
|
||||||
message = reaction.message
|
|
||||||
emoji = reaction.emoji
|
|
||||||
origid = self.wmessages.get(message.id, None)
|
|
||||||
if origid and reaction.count == 1:
|
|
||||||
cached_sent = self.message_cache.get(origid, ())
|
|
||||||
for _, msg in cached_sent:
|
|
||||||
# TODO: Would be better to have a Message and check the reactions
|
|
||||||
try:
|
|
||||||
if msg.id != message.id:
|
|
||||||
await msg.add_reaction(emoji)
|
|
||||||
except discord.HTTPException:
|
|
||||||
pass
|
|
||||||
|
|
||||||
async def fetch_webhook_for(self, channelid) -> discord.Webhook:
|
|
||||||
hook = self.hooks.get(channelid, None)
|
|
||||||
if hook is None:
|
|
||||||
row = await self.data.LinkHook.fetch(channelid)
|
|
||||||
if row is None:
|
|
||||||
channel = self.bot.get_channel(channelid)
|
|
||||||
if channel is None:
|
|
||||||
raise ValueError("Cannot find channel to create hook.")
|
|
||||||
hook = await channel.create_webhook(name="LabRat Channel Link")
|
|
||||||
await self.data.LinkHook.create(
|
|
||||||
channelid=channelid,
|
|
||||||
webhookid=hook.id,
|
|
||||||
token=hook.token,
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
hook = discord.Webhook.partial(row.webhookid, row.token, client=self.bot)
|
|
||||||
self.hooks[channelid] = hook
|
|
||||||
return hook
|
|
||||||
|
|
||||||
@cmds.hybrid_group(
|
|
||||||
name='linker',
|
|
||||||
description="Base command group for the channel linker"
|
|
||||||
)
|
|
||||||
@appcmds.default_permissions(manage_channels=True)
|
|
||||||
async def linker_group(self, ctx: LionContext):
|
|
||||||
...
|
|
||||||
|
|
||||||
@linker_group.command(
|
|
||||||
name='link',
|
|
||||||
description="Create a new link, or add a channel to an existing link."
|
|
||||||
)
|
|
||||||
@appcmds.describe(
|
|
||||||
name="Name of the new or existing channel link.",
|
|
||||||
channel1="First channel to add to the link.",
|
|
||||||
channel2="Second channel to add to the link.",
|
|
||||||
channel3="Third channel to add to the link.",
|
|
||||||
channel4="Fourth channel to add to the link.",
|
|
||||||
channel5="Fifth channel to add to the link.",
|
|
||||||
channelid="Optionally add a channel by id (for e.g. cross-server links).",
|
|
||||||
)
|
|
||||||
async def linker_link(self, ctx: LionContext,
|
|
||||||
name: str,
|
|
||||||
channel1: Optional[discord.TextChannel | discord.VoiceChannel] = None,
|
|
||||||
channel2: Optional[discord.TextChannel | discord.VoiceChannel] = None,
|
|
||||||
channel3: Optional[discord.TextChannel | discord.VoiceChannel] = None,
|
|
||||||
channel4: Optional[discord.TextChannel | discord.VoiceChannel] = None,
|
|
||||||
channel5: Optional[discord.TextChannel | discord.VoiceChannel] = None,
|
|
||||||
channelid: Optional[str] = None,
|
|
||||||
):
|
|
||||||
if not ctx.interaction:
|
|
||||||
return
|
|
||||||
await ctx.interaction.response.defer(thinking=True)
|
|
||||||
|
|
||||||
# Check if link 'name' already exists, create if not
|
|
||||||
existing = await self.data.Link.fetch_where()
|
|
||||||
link_row = next((row for row in existing if row.name.lower() == name.lower()), None)
|
|
||||||
if link_row is None:
|
|
||||||
# Create
|
|
||||||
link_row = await self.data.Link.create(name=name)
|
|
||||||
link_channels = set()
|
|
||||||
created = True
|
|
||||||
else:
|
|
||||||
records = await self.data.channel_links.select_where(linkid=link_row.linkid)
|
|
||||||
link_channels = {record['channelid'] for record in records}
|
|
||||||
created = False
|
|
||||||
|
|
||||||
# Create webhooks and webhook rows on channels if required
|
|
||||||
maybe_channels = [
|
|
||||||
channel1, channel2, channel3, channel4, channel5,
|
|
||||||
]
|
|
||||||
if channelid and channelid.isdigit():
|
|
||||||
channel = self.bot.get_channel(int(channelid))
|
|
||||||
maybe_channels.append(channel)
|
|
||||||
|
|
||||||
channels = [channel for channel in maybe_channels if channel]
|
|
||||||
for channel in channels:
|
|
||||||
await self.fetch_webhook_for(channel.id)
|
|
||||||
|
|
||||||
# Insert or update the links
|
|
||||||
for channel in channels:
|
|
||||||
if channel.id not in link_channels:
|
|
||||||
await self.data.channel_links.insert(linkid=link_row.linkid, channelid=channel.id)
|
|
||||||
|
|
||||||
await self.reload_links()
|
|
||||||
|
|
||||||
if created:
|
|
||||||
embed = discord.Embed(
|
|
||||||
colour=discord.Colour.brand_green(),
|
|
||||||
title="Link Created",
|
|
||||||
description=(
|
|
||||||
"Created the link **{name}** and linked channels:\n{channels}"
|
|
||||||
).format(name=name, channels=', '.join(channel.mention for channel in channels))
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
channelids = self.link_channels[link_row.linkid]
|
|
||||||
channelstr = ', '.join(f"<#{cid}>" for cid in channelids)
|
|
||||||
embed = discord.Embed(
|
|
||||||
colour=discord.Colour.brand_green(),
|
|
||||||
title="Channels Linked",
|
|
||||||
description=(
|
|
||||||
"Updated the link **{name}** to link the following channels:\n{channelstr}"
|
|
||||||
).format(name=link_row.name, channelstr=channelstr)
|
|
||||||
)
|
|
||||||
await ctx.reply(embed=embed)
|
|
||||||
|
|
||||||
@linker_group.command(
|
|
||||||
name='unlink',
|
|
||||||
description="Destroy a link, or remove a channel from a link."
|
|
||||||
)
|
|
||||||
@appcmds.describe(
|
|
||||||
name="Name of the link to destroy",
|
|
||||||
channel="Channel to remove from the link.",
|
|
||||||
)
|
|
||||||
async def linker_unlink(self, ctx: LionContext,
|
|
||||||
name: str, channel: Optional[GuildChannel] = None):
|
|
||||||
if not ctx.interaction:
|
|
||||||
return
|
|
||||||
# Get the link, error if it doesn't exist
|
|
||||||
existing = await self.data.Link.fetch_where()
|
|
||||||
link_row = next((row for row in existing if row.name.lower() == name.lower()), None)
|
|
||||||
if link_row is None:
|
|
||||||
raise UserInputError(
|
|
||||||
f"Link **{name}** doesn't exist!"
|
|
||||||
)
|
|
||||||
|
|
||||||
link_channelids = self.link_channels.get(link_row.linkid, ())
|
|
||||||
|
|
||||||
if channel is not None:
|
|
||||||
# If channel was given, remove channel from link and ack
|
|
||||||
if channel.id not in link_channelids:
|
|
||||||
raise UserInputError(
|
|
||||||
f"{channel.mention} is not linked in **{link_row.name}**!"
|
|
||||||
)
|
|
||||||
await self.data.channel_links.delete_where(channelid=channel.id, linkid=link_row.linkid)
|
|
||||||
embed = discord.Embed(
|
|
||||||
colour=discord.Colour.brand_green(),
|
|
||||||
title="Channel Unlinked",
|
|
||||||
description=f"{channel.mention} has been removed from **{link_row.name}**."
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
# Otherwise, confirm link destroy, delete link row, and ack
|
|
||||||
channels = ', '.join(f"<#{cid}>" for cid in link_channelids)
|
|
||||||
confirm = Confirm(
|
|
||||||
f"Are you sure you want to remove the link **{link_row.name}**?\nLinked channels: {channels}",
|
|
||||||
ctx.author.id,
|
|
||||||
)
|
|
||||||
confirm.embed.colour = discord.Colour.red()
|
|
||||||
try:
|
|
||||||
result = await confirm.ask(ctx.interaction)
|
|
||||||
except ResponseTimedOut:
|
|
||||||
result = False
|
|
||||||
if not result:
|
|
||||||
raise SafeCancellation
|
|
||||||
|
|
||||||
embed = discord.Embed(
|
|
||||||
colour=discord.Colour.brand_green(),
|
|
||||||
title="Link removed",
|
|
||||||
description=f"Link **{link_row.name}** removed, the following channels were unlinked:\n{channels}"
|
|
||||||
)
|
|
||||||
await link_row.delete()
|
|
||||||
|
|
||||||
await self.reload_links()
|
|
||||||
await ctx.reply(embed=embed)
|
|
||||||
|
|
||||||
@linker_link.autocomplete('name')
|
|
||||||
async def _acmpl_link_name(self, interaction: discord.Interaction, partial: str):
|
|
||||||
"""
|
|
||||||
Autocomplete an existing link.
|
|
||||||
"""
|
|
||||||
existing = await self.data.Link.fetch_where()
|
|
||||||
names = [row.name for row in existing]
|
|
||||||
matching = [row.name for row in existing if partial.lower() in row.name.lower()]
|
|
||||||
if not matching:
|
|
||||||
choice = appcmds.Choice(
|
|
||||||
name=f"Create a new link '{partial}'",
|
|
||||||
value=partial
|
|
||||||
)
|
|
||||||
choices = [choice]
|
|
||||||
else:
|
|
||||||
choices = [
|
|
||||||
appcmds.Choice(
|
|
||||||
name=f"Link {name}",
|
|
||||||
value=name
|
|
||||||
)
|
|
||||||
for name in matching
|
|
||||||
]
|
|
||||||
return choices
|
|
||||||
|
|
||||||
@linker_unlink.autocomplete('name')
|
|
||||||
async def _acmpl_unlink_name(self, interaction: discord.Interaction, partial: str):
|
|
||||||
"""
|
|
||||||
Autocomplete an existing link.
|
|
||||||
"""
|
|
||||||
existing = await self.data.Link.fetch_where()
|
|
||||||
matching = [row.name for row in existing if partial.lower() in row.name.lower()]
|
|
||||||
if not matching:
|
|
||||||
choice = appcmds.Choice(
|
|
||||||
name=f"No existing links matching '{partial}'",
|
|
||||||
value=partial
|
|
||||||
)
|
|
||||||
choices = [choice]
|
|
||||||
else:
|
|
||||||
choices = [
|
|
||||||
appcmds.Choice(
|
|
||||||
name=f"Link {name}",
|
|
||||||
value=name
|
|
||||||
)
|
|
||||||
for name in matching
|
|
||||||
]
|
|
||||||
return choices
|
|
||||||
|
|
||||||
@linker_group.command(
|
|
||||||
name='links',
|
|
||||||
description="Display the existing channel links."
|
|
||||||
)
|
|
||||||
async def linker_links(self, ctx: LionContext):
|
|
||||||
if not ctx.interaction:
|
|
||||||
return
|
|
||||||
await ctx.interaction.response.defer(thinking=True)
|
|
||||||
|
|
||||||
links = await self.data.Link.fetch_where()
|
|
||||||
|
|
||||||
if not links:
|
|
||||||
embed = discord.Embed(
|
|
||||||
colour=discord.Colour.light_grey(),
|
|
||||||
title="No channel links have been set up!",
|
|
||||||
description="Create a new link and add channels with {linker}".format(
|
|
||||||
linker=self.bot.core.mention_cmd('linker link')
|
|
||||||
)
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
embed = discord.Embed(
|
|
||||||
colour=discord.Colour.brand_green(),
|
|
||||||
title=f"Channel Links in {ctx.guild.name}",
|
|
||||||
)
|
|
||||||
for link in links:
|
|
||||||
channelids = self.link_channels.get(link.linkid, ())
|
|
||||||
channelstr = ', '.join(f"<#{cid}>" for cid in channelids)
|
|
||||||
embed.add_field(
|
|
||||||
name=f"Link **{link.name}**",
|
|
||||||
value=channelstr,
|
|
||||||
inline=False
|
|
||||||
)
|
|
||||||
# TODO: May want paging if over 25 links....
|
|
||||||
await ctx.reply(embed=embed)
|
|
||||||
|
|
||||||
@linker_group.command(
|
|
||||||
name="webhook",
|
|
||||||
description='Manually configure the webhook for a given channel.'
|
|
||||||
)
|
|
||||||
async def linker_webhook(self, ctx: LionContext, channel: discord.abc.GuildChannel, webhook: str):
|
|
||||||
if not ctx.interaction:
|
|
||||||
return
|
|
||||||
|
|
||||||
hook = discord.Webhook.from_url(webhook, client=self.bot)
|
|
||||||
existing = await self.data.LionHook.fetch(channel.id)
|
|
||||||
if existing:
|
|
||||||
await existing.update(webhookid=hook.id, token=hook.token)
|
|
||||||
else:
|
|
||||||
await self.data.LinkHook.create(
|
|
||||||
channelid=channel.id,
|
|
||||||
webhookid=hook.id,
|
|
||||||
token=hook.token,
|
|
||||||
)
|
|
||||||
self.hooks[channel.id] = hook
|
|
||||||
await ctx.reply(f"Webhook for {channel.mention} updated!")
|
|
||||||
@@ -1,39 +0,0 @@
|
|||||||
from data import Registry, RowModel, Table
|
|
||||||
from data.columns import Integer, Bool, Timestamp, String
|
|
||||||
|
|
||||||
|
|
||||||
class LinkData(Registry):
|
|
||||||
class Link(RowModel):
|
|
||||||
"""
|
|
||||||
Schema
|
|
||||||
------
|
|
||||||
CREATE TABLE links(
|
|
||||||
linkid SERIAL PRIMARY KEY,
|
|
||||||
name TEXT
|
|
||||||
);
|
|
||||||
"""
|
|
||||||
_tablename_ = 'links'
|
|
||||||
_cache_ = {}
|
|
||||||
|
|
||||||
linkid = Integer(primary=True)
|
|
||||||
name = String()
|
|
||||||
|
|
||||||
|
|
||||||
channel_links = Table('channel_links')
|
|
||||||
|
|
||||||
class LinkHook(RowModel):
|
|
||||||
"""
|
|
||||||
Schema
|
|
||||||
------
|
|
||||||
CREATE TABLE channel_webhooks(
|
|
||||||
channelid BIGINT PRIMARY KEY,
|
|
||||||
webhookid BIGINT NOT NULL,
|
|
||||||
token TEXT NOT NULL
|
|
||||||
);
|
|
||||||
"""
|
|
||||||
_tablename_ = 'channel_webhooks'
|
|
||||||
_cache_ = {}
|
|
||||||
|
|
||||||
channelid = Integer(primary=True)
|
|
||||||
webhookid = Integer()
|
|
||||||
token = String()
|
|
||||||
Submodule
+1
Submodule src/modules/voicelog added at 0cfc9b9986
@@ -2,6 +2,8 @@ import logging
|
|||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
async def setup(bot):
|
async def setup(bot):
|
||||||
from .cog import VoiceFixCog
|
from .cog import YarnCog
|
||||||
await bot.add_cog(VoiceFixCog(bot))
|
|
||||||
|
await bot.add_cog(YarnCog(bot))
|
||||||
@@ -0,0 +1,134 @@
|
|||||||
|
from typing import Literal
|
||||||
|
from collections import defaultdict
|
||||||
|
import datetime as dt
|
||||||
|
from datetime import datetime, timedelta, UTC
|
||||||
|
|
||||||
|
from data.queries import ORDER
|
||||||
|
import discord
|
||||||
|
from discord.ext import commands as cmds
|
||||||
|
from discord import app_commands as appcmds
|
||||||
|
|
||||||
|
from meta import LionBot, LionCog, LionContext
|
||||||
|
from meta.logger import log_wrap
|
||||||
|
from utils.lib import strfdur, utc_now, strfdur, paginate_list, pager
|
||||||
|
|
||||||
|
from modules.voicelog.plugin.data import VoiceLogSession
|
||||||
|
|
||||||
|
|
||||||
|
class YarnCog(LionCog):
|
||||||
|
"""
|
||||||
|
Assorted toys for Lilac
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, bot: LionBot):
|
||||||
|
self.bot = bot
|
||||||
|
self.desired_voice = defaultdict(dict)
|
||||||
|
|
||||||
|
@LionCog.listener("on_voice_state_update")
|
||||||
|
async def voicestate_muter(self, member, before, after):
|
||||||
|
if not after.channel:
|
||||||
|
return
|
||||||
|
target_state = self.desired_voice[member.guild.id].pop(member.id, None)
|
||||||
|
if target_state is None:
|
||||||
|
return
|
||||||
|
await member.edit(mute=target_state)
|
||||||
|
# TODO: Log using voicelog webhook
|
||||||
|
|
||||||
|
@LionCog.listener("on_reaction_add")
|
||||||
|
async def lilac_confirms(self, reaction: discord.Reaction, user: discord.User):
|
||||||
|
if not reaction.me:
|
||||||
|
await reaction.message.add_reaction(reaction.emoji)
|
||||||
|
|
||||||
|
@LionCog.listener("on_reaction_remove")
|
||||||
|
async def lilac_unconfirms(self, reaction: discord.Reaction, user: discord.User):
|
||||||
|
if (
|
||||||
|
reaction.me
|
||||||
|
and reaction.count == 1
|
||||||
|
and reaction.message.author.id != self.bot.user.id
|
||||||
|
):
|
||||||
|
await reaction.remove(self.bot.user)
|
||||||
|
|
||||||
|
@cmds.hybrid_command(name="voicestate")
|
||||||
|
@cmds.has_guild_permissions(mute_members=True)
|
||||||
|
async def voicestate_cmd(
|
||||||
|
self, ctx, user: discord.Member, state: Literal["muted", "unmuted", "clear"]
|
||||||
|
):
|
||||||
|
self.desired_voice[ctx.guild.id].pop(user.id, None)
|
||||||
|
|
||||||
|
if state == "clear":
|
||||||
|
# We've already removed the saved state, don't do anything else.
|
||||||
|
ack = f"{user.mention} target voice state cleared!"
|
||||||
|
elif user.voice:
|
||||||
|
# If user is currently in channel, apply the state
|
||||||
|
if state == "muted":
|
||||||
|
await user.edit(mute=True)
|
||||||
|
ack = f"{user.mention} muted!"
|
||||||
|
elif state == "unmuted":
|
||||||
|
await user.edit(mute=False)
|
||||||
|
ack = f"{user.mention} unmuted!"
|
||||||
|
else:
|
||||||
|
# If user is not currently in channel, save the state
|
||||||
|
if state == "muted":
|
||||||
|
self.desired_voice[ctx.guild.id][user.id] = True
|
||||||
|
ack = f"{user.mention} will be muted!"
|
||||||
|
elif state == "unmuted":
|
||||||
|
self.desired_voice[ctx.guild.id][user.id] = False
|
||||||
|
ack = f"{user.mention} will be unmuted!"
|
||||||
|
await ctx.reply(
|
||||||
|
embed=discord.Embed(colour=discord.Colour.brand_green(), description=ack)
|
||||||
|
)
|
||||||
|
|
||||||
|
@cmds.hybrid_command(name="topvoice")
|
||||||
|
async def topvoice_cmd(self, ctx):
|
||||||
|
"""
|
||||||
|
Show top voice members by total time.
|
||||||
|
"""
|
||||||
|
target_channelid = 1383707078740279366
|
||||||
|
since_stamp = 1769832959
|
||||||
|
|
||||||
|
voicelogger = ctx.bot.get_cog("VoiceLogCog")
|
||||||
|
session_data = voicelogger.data.voicelog_sessions
|
||||||
|
|
||||||
|
query = (
|
||||||
|
session_data.select_where(
|
||||||
|
VoiceLogSession.joined_at
|
||||||
|
>= datetime.fromtimestamp(since_stamp, tz=UTC),
|
||||||
|
guildid=ctx.guild.id,
|
||||||
|
channelid=target_channelid,
|
||||||
|
)
|
||||||
|
.select(
|
||||||
|
userid="userid",
|
||||||
|
total_time="SUM(COALESCE(duration, EXTRACT(EPOCH FROM (NOW() - joined_at))))",
|
||||||
|
)
|
||||||
|
.order_by("total_time", ORDER.DESC)
|
||||||
|
.group_by("userid")
|
||||||
|
.with_no_adapter()
|
||||||
|
)
|
||||||
|
leaderboard = [(row["userid"], int(row["total_time"])) for row in await query]
|
||||||
|
|
||||||
|
# Format for display and pager
|
||||||
|
# First collect names
|
||||||
|
names = {}
|
||||||
|
for uid, _ in leaderboard:
|
||||||
|
user = ctx.guild.get_member(uid)
|
||||||
|
if user is None:
|
||||||
|
try:
|
||||||
|
user = await ctx.guild.fetch_member(uid)
|
||||||
|
except discord.NotFound:
|
||||||
|
user = None
|
||||||
|
names[uid] = user.display_name if user else str(uid)
|
||||||
|
|
||||||
|
lb_strings = []
|
||||||
|
max_name_len = min((30, max(len(name) for name in names.values())))
|
||||||
|
for i, (uid, total) in enumerate(leaderboard):
|
||||||
|
lb_strings.append(
|
||||||
|
"{:<{}}\t{:<9}".format(
|
||||||
|
names[uid], max_name_len, strfdur(total, short=False)
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
page_len = 20
|
||||||
|
title = "Voice Leaderboard"
|
||||||
|
pages = paginate_list(lb_strings, block_length=page_len, title=title)
|
||||||
|
|
||||||
|
await pager(ctx, pages)
|
||||||
@@ -0,0 +1,7 @@
|
|||||||
|
from babel.translator import LocalBabel
|
||||||
|
babel = LocalBabel('settings_base')
|
||||||
|
|
||||||
|
from .data import ModelData, ListData
|
||||||
|
from .base import BaseSetting
|
||||||
|
from .ui import SettingWidget, InteractiveSetting
|
||||||
|
from .groups import SettingDotDict, SettingGroup, ModelSettings, ModelSetting
|
||||||
@@ -0,0 +1,166 @@
|
|||||||
|
from typing import Generic, TypeVar, Type, Optional, overload
|
||||||
|
|
||||||
|
|
||||||
|
"""
|
||||||
|
Setting metclass?
|
||||||
|
Parse setting docstring to generate default info?
|
||||||
|
Or just put it in the decorator we are already using
|
||||||
|
"""
|
||||||
|
|
||||||
|
|
||||||
|
# Typing using Generic[parent_id_type, data_type, value_type]
|
||||||
|
# value generic, could be Union[?, UNSET]
|
||||||
|
ParentID = TypeVar('ParentID')
|
||||||
|
SettingData = TypeVar('SettingData')
|
||||||
|
SettingValue = TypeVar('SettingValue')
|
||||||
|
|
||||||
|
T = TypeVar('T', bound='BaseSetting')
|
||||||
|
|
||||||
|
|
||||||
|
class BaseSetting(Generic[ParentID, SettingData, SettingValue]):
|
||||||
|
"""
|
||||||
|
Abstract base class describing a stored configuration setting.
|
||||||
|
A setting consists of logic to load the setting from storage,
|
||||||
|
present it in a readable form, understand user entered values,
|
||||||
|
and write it again in storage.
|
||||||
|
Additionally, the setting has attributes attached describing
|
||||||
|
the setting in a user-friendly manner for display purposes.
|
||||||
|
"""
|
||||||
|
setting_id: str # Unique source identifier for the setting
|
||||||
|
|
||||||
|
_default: Optional[SettingData] = None # Default data value for the setting
|
||||||
|
|
||||||
|
def __init__(self, parent_id: ParentID, data: Optional[SettingData], **kwargs):
|
||||||
|
self.parent_id = parent_id
|
||||||
|
self._data = data
|
||||||
|
self.kwargs = kwargs
|
||||||
|
|
||||||
|
# Instance generation
|
||||||
|
@classmethod
|
||||||
|
async def get(cls: Type[T], parent_id: ParentID, **kwargs) -> T:
|
||||||
|
"""
|
||||||
|
Return a setting instance initialised from the stored value, associated with the given parent id.
|
||||||
|
"""
|
||||||
|
data = await cls._reader(parent_id, **kwargs)
|
||||||
|
return cls(parent_id, data, **kwargs)
|
||||||
|
|
||||||
|
# Main interface
|
||||||
|
@property
|
||||||
|
def data(self) -> Optional[SettingData]:
|
||||||
|
"""
|
||||||
|
Retrieves the current internal setting data if it is set, otherwise the default data
|
||||||
|
"""
|
||||||
|
return self._data if self._data is not None else self.default
|
||||||
|
|
||||||
|
@data.setter
|
||||||
|
def data(self, new_data: Optional[SettingData]):
|
||||||
|
"""
|
||||||
|
Sets the internal raw data.
|
||||||
|
Does not write the changes.
|
||||||
|
"""
|
||||||
|
self._data = new_data
|
||||||
|
|
||||||
|
@property
|
||||||
|
def default(self) -> Optional[SettingData]:
|
||||||
|
"""
|
||||||
|
Retrieves the default value for this setting.
|
||||||
|
Settings should override this if the default depends on the object id.
|
||||||
|
"""
|
||||||
|
return self._default
|
||||||
|
|
||||||
|
@property
|
||||||
|
def value(self) -> SettingValue: # Actually optional *if* _default is None
|
||||||
|
"""
|
||||||
|
Context-aware object or objects associated with the setting.
|
||||||
|
"""
|
||||||
|
return self._data_to_value(self.parent_id, self.data) # type: ignore
|
||||||
|
|
||||||
|
@value.setter
|
||||||
|
def value(self, new_value: Optional[SettingValue]):
|
||||||
|
"""
|
||||||
|
Setter which reads the discord-aware object and converts it to data.
|
||||||
|
Does not write the new value.
|
||||||
|
"""
|
||||||
|
self._data = self._data_from_value(self.parent_id, new_value)
|
||||||
|
|
||||||
|
async def write(self, **kwargs) -> None:
|
||||||
|
"""
|
||||||
|
Write current data to the database.
|
||||||
|
For settings which override this,
|
||||||
|
ensure you handle deletion of values when internal data is None.
|
||||||
|
"""
|
||||||
|
await self._writer(self.parent_id, self._data, **kwargs)
|
||||||
|
|
||||||
|
# Raw converters
|
||||||
|
@overload
|
||||||
|
@classmethod
|
||||||
|
def _data_from_value(cls: Type[T], parent_id: ParentID, value: SettingValue, **kwargs) -> SettingData:
|
||||||
|
...
|
||||||
|
|
||||||
|
@overload
|
||||||
|
@classmethod
|
||||||
|
def _data_from_value(cls: Type[T], parent_id: ParentID, value: None, **kwargs) -> None:
|
||||||
|
...
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _data_from_value(
|
||||||
|
cls: Type[T], parent_id: ParentID, value: Optional[SettingValue], **kwargs
|
||||||
|
) -> Optional[SettingData]:
|
||||||
|
"""
|
||||||
|
Convert a high-level setting value to internal data.
|
||||||
|
Must be overridden by the setting.
|
||||||
|
Be aware of UNSET values, these should always pass through as None
|
||||||
|
to provide an unsetting interface.
|
||||||
|
"""
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
@overload
|
||||||
|
@classmethod
|
||||||
|
def _data_to_value(cls: Type[T], parent_id: ParentID, data: SettingData, **kwargs) -> SettingValue:
|
||||||
|
...
|
||||||
|
|
||||||
|
@overload
|
||||||
|
@classmethod
|
||||||
|
def _data_to_value(cls: Type[T], parent_id: ParentID, data: None, **kwargs) -> None:
|
||||||
|
...
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _data_to_value(
|
||||||
|
cls: Type[T], parent_id: ParentID, data: Optional[SettingData], **kwargs
|
||||||
|
) -> Optional[SettingValue]:
|
||||||
|
"""
|
||||||
|
Convert internal data to high-level setting value.
|
||||||
|
Must be overriden by the setting.
|
||||||
|
"""
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
# Database access
|
||||||
|
@classmethod
|
||||||
|
async def _reader(cls: Type[T], parent_id: ParentID, **kwargs) -> Optional[SettingData]:
|
||||||
|
"""
|
||||||
|
Retrieve the setting data associated with the given parent_id.
|
||||||
|
May be None if the setting is not set.
|
||||||
|
Must be overridden by the setting.
|
||||||
|
"""
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
async def _writer(cls: Type[T], parent_id: ParentID, data: Optional[SettingData], **kwargs) -> None:
|
||||||
|
"""
|
||||||
|
Write provided setting data to storage.
|
||||||
|
Must be overridden by the setting unless the `write` method is overridden.
|
||||||
|
If the data is None, the setting is UNSET and should be deleted.
|
||||||
|
"""
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
async def setup(cls, bot):
|
||||||
|
"""
|
||||||
|
Initialisation task to be executed during client initialisation.
|
||||||
|
May be used for e.g. populating a cache or required client setup.
|
||||||
|
|
||||||
|
Main application must execute the initialisation task before the setting is used.
|
||||||
|
Further, the task must always be executable, if the setting is loaded.
|
||||||
|
Conditional initialisation should go in the relevant module's init tasks.
|
||||||
|
"""
|
||||||
|
return None
|
||||||
@@ -0,0 +1,233 @@
|
|||||||
|
from typing import Type
|
||||||
|
import json
|
||||||
|
|
||||||
|
from data import RowModel, Table, ORDER
|
||||||
|
from meta.logger import log_wrap, set_logging_context
|
||||||
|
|
||||||
|
|
||||||
|
class ModelData:
|
||||||
|
"""
|
||||||
|
Mixin for settings stored in a single row and column of a Model.
|
||||||
|
Assumes that the parent_id is the identity key of the Model.
|
||||||
|
|
||||||
|
This does not create a reference to the Row.
|
||||||
|
"""
|
||||||
|
# Table storing the desired data
|
||||||
|
_model: Type[RowModel]
|
||||||
|
|
||||||
|
# Column with the desired data
|
||||||
|
_column: str
|
||||||
|
|
||||||
|
# Whether to create a row if not found
|
||||||
|
_create_row = False
|
||||||
|
|
||||||
|
# High level data cache to use, leave as None to disable cache.
|
||||||
|
_cache = None # Map[id -> value]
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _read_from_row(cls, parent_id, row, **kwargs):
|
||||||
|
data = row[cls._column]
|
||||||
|
|
||||||
|
if cls._cache is not None:
|
||||||
|
cls._cache[parent_id] = data
|
||||||
|
|
||||||
|
return data
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
async def _reader(cls, parent_id, use_cache=True, **kwargs):
|
||||||
|
"""
|
||||||
|
Read in the requested column associated to the parent id.
|
||||||
|
"""
|
||||||
|
if cls._cache is not None and parent_id in cls._cache and use_cache:
|
||||||
|
return cls._cache[parent_id]
|
||||||
|
|
||||||
|
model = cls._model
|
||||||
|
if cls._create_row:
|
||||||
|
row = await model.fetch_or_create(parent_id)
|
||||||
|
else:
|
||||||
|
row = await model.fetch(parent_id)
|
||||||
|
data = row[cls._column] if row else None
|
||||||
|
|
||||||
|
if cls._cache is not None:
|
||||||
|
cls._cache[parent_id] = data
|
||||||
|
|
||||||
|
return data
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
async def _writer(cls, parent_id, data, **kwargs):
|
||||||
|
"""
|
||||||
|
Write the provided entry to the table.
|
||||||
|
This does *not* create the row if it does not exist.
|
||||||
|
It only updates.
|
||||||
|
"""
|
||||||
|
# TODO: Better way of getting the key?
|
||||||
|
# TODO: Transaction
|
||||||
|
if not isinstance(parent_id, tuple):
|
||||||
|
parent_id = (parent_id, )
|
||||||
|
model = cls._model
|
||||||
|
rows = await model.table.update_where(
|
||||||
|
**model._dict_from_id(parent_id)
|
||||||
|
).set(
|
||||||
|
**{cls._column: data}
|
||||||
|
)
|
||||||
|
# If we didn't update any rows, create a new row
|
||||||
|
if not rows:
|
||||||
|
await model.fetch_or_create(**model._dict_from_id(parent_id), **{cls._column: data})
|
||||||
|
|
||||||
|
if cls._cache is not None:
|
||||||
|
cls._cache[parent_id] = data
|
||||||
|
|
||||||
|
|
||||||
|
class ListData:
|
||||||
|
"""
|
||||||
|
Mixin for list types implemented on a Table.
|
||||||
|
Implements a reader and writer.
|
||||||
|
This assumes the list is the only data stored in the table,
|
||||||
|
and removes list entries by deleting rows.
|
||||||
|
"""
|
||||||
|
setting_id: str
|
||||||
|
|
||||||
|
# Table storing the setting data
|
||||||
|
_table_interface: Table
|
||||||
|
|
||||||
|
# Name of the column storing the id
|
||||||
|
_id_column: str
|
||||||
|
|
||||||
|
# Name of the column storing the data to read
|
||||||
|
_data_column: str
|
||||||
|
|
||||||
|
# Name of column storing the order index to use, if any. Assumed to be Serial on writing.
|
||||||
|
_order_column: str
|
||||||
|
_order_type: ORDER = ORDER.ASC
|
||||||
|
|
||||||
|
# High level data cache to use, set to None to disable cache.
|
||||||
|
_cache = None # Map[id -> value]
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
@log_wrap(isolate=True)
|
||||||
|
async def _reader(cls, parent_id, use_cache=True, **kwargs):
|
||||||
|
"""
|
||||||
|
Read in all entries associated to the given id.
|
||||||
|
"""
|
||||||
|
set_logging_context(action="Read cls.setting_id")
|
||||||
|
if cls._cache is not None and parent_id in cls._cache and use_cache:
|
||||||
|
return cls._cache[parent_id]
|
||||||
|
|
||||||
|
table = cls._table_interface # type: Table
|
||||||
|
query = table.select_where(**{cls._id_column: parent_id}).select(cls._data_column)
|
||||||
|
if cls._order_column:
|
||||||
|
query.order_by(cls._order_column, direction=cls._order_type)
|
||||||
|
|
||||||
|
rows = await query
|
||||||
|
data = [row[cls._data_column] for row in rows]
|
||||||
|
|
||||||
|
if cls._cache is not None:
|
||||||
|
cls._cache[parent_id] = data
|
||||||
|
|
||||||
|
return data
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
@log_wrap(isolate=True)
|
||||||
|
async def _writer(cls, id, data, add_only=False, remove_only=False, **kwargs):
|
||||||
|
"""
|
||||||
|
Write the provided list to storage.
|
||||||
|
"""
|
||||||
|
set_logging_context(action="Write cls.setting_id")
|
||||||
|
table = cls._table_interface
|
||||||
|
async with table.connector.connection() as conn:
|
||||||
|
table.connector.conn = conn
|
||||||
|
async with conn.transaction():
|
||||||
|
# Handle None input as an empty list
|
||||||
|
if data is None:
|
||||||
|
data = []
|
||||||
|
|
||||||
|
current = await cls._reader(id, use_cache=False, **kwargs)
|
||||||
|
if not cls._order_column and (add_only or remove_only):
|
||||||
|
to_insert = [item for item in data if item not in current] if not remove_only else []
|
||||||
|
to_remove = data if remove_only else (
|
||||||
|
[item for item in current if item not in data] if not add_only else []
|
||||||
|
)
|
||||||
|
|
||||||
|
# Handle required deletions
|
||||||
|
if to_remove:
|
||||||
|
params = {
|
||||||
|
cls._id_column: id,
|
||||||
|
cls._data_column: to_remove
|
||||||
|
}
|
||||||
|
await table.delete_where(**params)
|
||||||
|
|
||||||
|
# Handle required insertions
|
||||||
|
if to_insert:
|
||||||
|
columns = (cls._id_column, cls._data_column)
|
||||||
|
values = [(id, value) for value in to_insert]
|
||||||
|
await table.insert_many(columns, *values)
|
||||||
|
|
||||||
|
if cls._cache is not None:
|
||||||
|
new_current = [item for item in current + to_insert if item not in to_remove]
|
||||||
|
cls._cache[id] = new_current
|
||||||
|
else:
|
||||||
|
# Remove all and add all to preserve order
|
||||||
|
delete_params = {cls._id_column: id}
|
||||||
|
await table.delete_where(**delete_params)
|
||||||
|
|
||||||
|
if data:
|
||||||
|
columns = (cls._id_column, cls._data_column)
|
||||||
|
values = [(id, value) for value in data]
|
||||||
|
await table.insert_many(columns, *values)
|
||||||
|
|
||||||
|
if cls._cache is not None:
|
||||||
|
cls._cache[id] = data
|
||||||
|
|
||||||
|
|
||||||
|
class KeyValueData:
|
||||||
|
"""
|
||||||
|
Mixin for settings implemented in a Key-Value table.
|
||||||
|
The underlying table should have a Unique constraint on the `(_id_column, _key_column)` pair.
|
||||||
|
"""
|
||||||
|
_table_interface: Table
|
||||||
|
|
||||||
|
_id_column: str
|
||||||
|
|
||||||
|
_key_column: str
|
||||||
|
|
||||||
|
_value_column: str
|
||||||
|
|
||||||
|
_key: str
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
async def _reader(cls, id, **kwargs):
|
||||||
|
params = {
|
||||||
|
cls._id_column: id,
|
||||||
|
cls._key_column: cls._key
|
||||||
|
}
|
||||||
|
|
||||||
|
row = await cls._table_interface.select_one_where(**params).select(cls._value_column)
|
||||||
|
data = row[cls._value_column] if row else None
|
||||||
|
|
||||||
|
if data is not None:
|
||||||
|
data = json.loads(data)
|
||||||
|
|
||||||
|
return data
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
async def _writer(cls, id, data, **kwargs):
|
||||||
|
params = {
|
||||||
|
cls._id_column: id,
|
||||||
|
cls._key_column: cls._key
|
||||||
|
}
|
||||||
|
if data is not None:
|
||||||
|
values = {
|
||||||
|
cls._value_column: json.dumps(data)
|
||||||
|
}
|
||||||
|
rows = await cls._table_interface.update_where(**params).set(**values)
|
||||||
|
if not rows:
|
||||||
|
await cls._table_interface.insert_many(
|
||||||
|
(cls._id_column, cls._key_column, cls._value_column),
|
||||||
|
(id, cls._key, json.dumps(data))
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
await cls._table_interface.delete_where(**params)
|
||||||
|
|
||||||
|
|
||||||
|
# class UserInputError(SafeCancellation):
|
||||||
|
# pass
|
||||||
@@ -0,0 +1,204 @@
|
|||||||
|
from typing import Generic, Type, TypeVar, Optional, overload
|
||||||
|
|
||||||
|
from data import RowModel
|
||||||
|
|
||||||
|
from .data import ModelData
|
||||||
|
from .ui import InteractiveSetting
|
||||||
|
from .base import BaseSetting
|
||||||
|
|
||||||
|
from utils.lib import tabulate
|
||||||
|
|
||||||
|
|
||||||
|
T = TypeVar('T', bound=InteractiveSetting)
|
||||||
|
|
||||||
|
|
||||||
|
class SettingDotDict(Generic[T], dict[str, Type[T]]):
|
||||||
|
"""
|
||||||
|
Dictionary structure allowing simple dot access to items.
|
||||||
|
"""
|
||||||
|
__getattr__ = dict.__getitem__ # type: ignore
|
||||||
|
__setattr__ = dict.__setitem__ # type: ignore
|
||||||
|
__delattr__ = dict.__delitem__ # type: ignore
|
||||||
|
|
||||||
|
|
||||||
|
class SettingGroup:
|
||||||
|
"""
|
||||||
|
A SettingGroup is a collection of settings under one name.
|
||||||
|
"""
|
||||||
|
__initial_settings__: list[Type[InteractiveSetting]] = []
|
||||||
|
|
||||||
|
_title: Optional[str] = None
|
||||||
|
_description: Optional[str] = None
|
||||||
|
|
||||||
|
def __init_subclass__(cls, title: Optional[str] = None):
|
||||||
|
cls._title = title or cls._title
|
||||||
|
cls._description = cls._description or cls.__doc__
|
||||||
|
|
||||||
|
settings: list[Type[InteractiveSetting]] = []
|
||||||
|
for item in cls.__dict__.values():
|
||||||
|
if isinstance(item, type) and issubclass(item, InteractiveSetting):
|
||||||
|
settings.append(item)
|
||||||
|
cls.__initial_settings__ = settings
|
||||||
|
|
||||||
|
def __init_settings__(self):
|
||||||
|
settings = SettingDotDict()
|
||||||
|
for setting in self.__initial_settings__:
|
||||||
|
settings[setting.__name__] = setting
|
||||||
|
return settings
|
||||||
|
|
||||||
|
def __init__(self, title=None, description=None) -> None:
|
||||||
|
self.title: str = title or self._title or self.__class__.__name__
|
||||||
|
self.description: str = description or self._description or ""
|
||||||
|
self.settings: SettingDotDict[InteractiveSetting] = self.__init_settings__()
|
||||||
|
|
||||||
|
def attach(self, cls: Type[T], name: Optional[str] = None):
|
||||||
|
name = name or cls.setting_id
|
||||||
|
self.settings[name] = cls
|
||||||
|
return cls
|
||||||
|
|
||||||
|
def detach(self, cls):
|
||||||
|
return self.settings.pop(cls.__name__, None)
|
||||||
|
|
||||||
|
def update(self, smap):
|
||||||
|
self.settings.update(smap.settings)
|
||||||
|
|
||||||
|
def reduce(self, *keys):
|
||||||
|
for key in keys:
|
||||||
|
self.settings.pop(key, None)
|
||||||
|
return
|
||||||
|
|
||||||
|
async def make_setting_table(self, parent_id, **kwargs):
|
||||||
|
"""
|
||||||
|
Convenience method for generating a rendered setting table.
|
||||||
|
"""
|
||||||
|
rows = []
|
||||||
|
for setting in self.settings.values():
|
||||||
|
if not setting._virtual:
|
||||||
|
set = await setting.get(parent_id, **kwargs)
|
||||||
|
name = set.display_name
|
||||||
|
value = str(set.formatted)
|
||||||
|
rows.append((name, value, set.hover_desc))
|
||||||
|
table_rows = tabulate(
|
||||||
|
*rows,
|
||||||
|
row_format="[`{invis}{key:<{pad}}{colon}`](https://lionbot.org \"{field[2]}\")\t{value}"
|
||||||
|
)
|
||||||
|
return '\n'.join(table_rows)
|
||||||
|
|
||||||
|
|
||||||
|
class ModelSetting(ModelData, BaseSetting):
|
||||||
|
...
|
||||||
|
|
||||||
|
|
||||||
|
class ModelConfig:
|
||||||
|
"""
|
||||||
|
A ModelConfig provides a central point of configuration for any object described by a single Model.
|
||||||
|
|
||||||
|
An instance of a ModelConfig represents configuration for a single object
|
||||||
|
(given by a single row of the corresponding Model).
|
||||||
|
|
||||||
|
The ModelConfig also supports registration of non-model configuration,
|
||||||
|
to support associated settings (e.g. list-settings) for the object.
|
||||||
|
|
||||||
|
This is an ABC, and must be subclassed for each object-type.
|
||||||
|
"""
|
||||||
|
settings: SettingDotDict
|
||||||
|
_model_settings: set
|
||||||
|
model: Type[RowModel]
|
||||||
|
|
||||||
|
def __init__(self, parent_id, row, **kwargs):
|
||||||
|
self.parent_id = parent_id
|
||||||
|
self.row = row
|
||||||
|
self.kwargs = kwargs
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def register_setting(cls, setting_cls):
|
||||||
|
"""
|
||||||
|
Decorator to register a non-model setting as part of the object configuration.
|
||||||
|
|
||||||
|
The setting class may be re-accessed through the `settings` class attr.
|
||||||
|
|
||||||
|
Subclasses may provide alternative access pathways to key non-model settings.
|
||||||
|
"""
|
||||||
|
cls.settings[setting_cls.setting_id] = setting_cls
|
||||||
|
return setting_cls
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def register_model_setting(cls, model_setting_cls):
|
||||||
|
"""
|
||||||
|
Decorator to register a model setting as part of the object configuration.
|
||||||
|
|
||||||
|
The setting class may be accessed through the `settings` class attr.
|
||||||
|
|
||||||
|
A fresh setting instance may also be retrieved (using cached data)
|
||||||
|
through the `get` instance method.
|
||||||
|
|
||||||
|
Subclasses are recommended to provide model settings as properties
|
||||||
|
for simplified access and type checking.
|
||||||
|
"""
|
||||||
|
cls._model_settings.add(model_setting_cls.setting_id)
|
||||||
|
return cls.register_setting(model_setting_cls)
|
||||||
|
|
||||||
|
def get(self, setting_id):
|
||||||
|
"""
|
||||||
|
Retrieve a freshly initialised copy of the given model-setting.
|
||||||
|
|
||||||
|
The given `setting_id` must have been previously registered through `register_model_setting`.
|
||||||
|
This uses cached data, and so is not guaranteed to be up-to-date.
|
||||||
|
"""
|
||||||
|
if setting_id not in self._model_settings:
|
||||||
|
# TODO: Log
|
||||||
|
raise ValueError
|
||||||
|
setting_cls = self.settings[setting_id]
|
||||||
|
data = setting_cls._read_from_row(self.parent_id, self.row, **self.kwargs)
|
||||||
|
return setting_cls(self.parent_id, data, **self.kwargs)
|
||||||
|
|
||||||
|
|
||||||
|
class ModelSettings:
|
||||||
|
"""
|
||||||
|
A ModelSettings instance aggregates multiple `ModelSetting` instances
|
||||||
|
bound to the same parent id on a single Model.
|
||||||
|
|
||||||
|
This enables a single point of access
|
||||||
|
for settings of a given Model,
|
||||||
|
with support for caching or deriving as needed.
|
||||||
|
|
||||||
|
This is an abstract base class,
|
||||||
|
and should be subclassed to define the contained settings.
|
||||||
|
"""
|
||||||
|
_settings: SettingDotDict = SettingDotDict()
|
||||||
|
model: Type[RowModel]
|
||||||
|
|
||||||
|
def __init__(self, parent_id, row, **kwargs):
|
||||||
|
self.parent_id = parent_id
|
||||||
|
self.row = row
|
||||||
|
self.kwargs = kwargs
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
async def fetch(cls, *parent_id, **kwargs):
|
||||||
|
"""
|
||||||
|
Load an instance of this ModelSetting with the given parent_id
|
||||||
|
and setting keyword arguments.
|
||||||
|
"""
|
||||||
|
row = await cls.model.fetch_or_create(*parent_id)
|
||||||
|
return cls(parent_id, row, **kwargs)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def attach(self, setting_cls):
|
||||||
|
"""
|
||||||
|
Decorator to attach the given setting class to this modelsetting.
|
||||||
|
"""
|
||||||
|
# This violates the interface principle, use structured typing instead?
|
||||||
|
if not (issubclass(setting_cls, BaseSetting) and issubclass(setting_cls, ModelData)):
|
||||||
|
raise ValueError(
|
||||||
|
f"The provided setting class must be `ModelSetting`, not {setting_cls.__class__.__name__}."
|
||||||
|
)
|
||||||
|
self._settings[setting_cls.setting_id] = setting_cls
|
||||||
|
return setting_cls
|
||||||
|
|
||||||
|
def get(self, setting_id):
|
||||||
|
setting_cls = self._settings.get(setting_id)
|
||||||
|
data = setting_cls._read_from_row(self.parent_id, self.row, **self.kwargs)
|
||||||
|
return setting_cls(self.parent_id, data, **self.kwargs)
|
||||||
|
|
||||||
|
def __getitem__(self, setting_id):
|
||||||
|
return self.get(setting_id)
|
||||||
@@ -0,0 +1,13 @@
|
|||||||
|
import discord
|
||||||
|
from discord import app_commands
|
||||||
|
|
||||||
|
|
||||||
|
class LocalString:
|
||||||
|
def __init__(self, string):
|
||||||
|
self.string = string
|
||||||
|
|
||||||
|
def as_string(self):
|
||||||
|
return self.string
|
||||||
|
|
||||||
|
|
||||||
|
_ = LocalString
|
||||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,512 @@
|
|||||||
|
from typing import Optional, Callable, Any, Dict, Coroutine, Generic, TypeVar, List
|
||||||
|
import asyncio
|
||||||
|
from contextvars import copy_context
|
||||||
|
|
||||||
|
import discord
|
||||||
|
from discord import ui
|
||||||
|
from discord.ui.button import ButtonStyle, Button, button
|
||||||
|
from discord.ui.modal import Modal
|
||||||
|
from discord.ui.text_input import TextInput
|
||||||
|
from meta.errors import UserInputError
|
||||||
|
|
||||||
|
from utils.lib import tabulate, recover_context
|
||||||
|
from utils.ui import FastModal
|
||||||
|
from meta.config import conf
|
||||||
|
from meta.context import ctx_bot
|
||||||
|
from babel.translator import ctx_translator, LazyStr
|
||||||
|
|
||||||
|
from .base import BaseSetting, ParentID, SettingData, SettingValue
|
||||||
|
from . import babel
|
||||||
|
|
||||||
|
_p = babel._p
|
||||||
|
|
||||||
|
|
||||||
|
ST = TypeVar('ST', bound='InteractiveSetting')
|
||||||
|
|
||||||
|
|
||||||
|
class SettingModal(FastModal):
|
||||||
|
input_field: TextInput = TextInput(label="Edit Setting")
|
||||||
|
|
||||||
|
def update_field(self, new_field):
|
||||||
|
self.remove_item(self.input_field)
|
||||||
|
self.add_item(new_field)
|
||||||
|
self.input_field = new_field
|
||||||
|
|
||||||
|
|
||||||
|
class SettingWidget(Generic[ST], ui.View):
|
||||||
|
# TODO: Permission restrictions and callback!
|
||||||
|
# Context variables for permitted user(s)? Subclass ui.View with PermittedView?
|
||||||
|
# Don't need to descend permissions to Modal
|
||||||
|
# Maybe combine with timeout manager
|
||||||
|
|
||||||
|
def __init__(self, setting: ST, auto_write=True, **kwargs):
|
||||||
|
self.setting = setting
|
||||||
|
self.update_children()
|
||||||
|
super().__init__(**kwargs)
|
||||||
|
self.auto_write = auto_write
|
||||||
|
|
||||||
|
self._interaction: Optional[discord.Interaction] = None
|
||||||
|
self._modal: Optional[SettingModal] = None
|
||||||
|
self._exports: List[ui.Item] = self.make_exports()
|
||||||
|
|
||||||
|
self._context = copy_context()
|
||||||
|
|
||||||
|
def update_children(self):
|
||||||
|
"""
|
||||||
|
Method called before base View initialisation.
|
||||||
|
Allows updating the children components (usually explicitly defined callbacks),
|
||||||
|
before Item instantiation.
|
||||||
|
"""
|
||||||
|
pass
|
||||||
|
|
||||||
|
def order_children(self, *children):
|
||||||
|
"""
|
||||||
|
Helper method to set and order the children using bound methods.
|
||||||
|
"""
|
||||||
|
child_map = {child.__name__: child for child in self.__view_children_items__}
|
||||||
|
self.__view_children_items__ = [child_map[child.__name__] for child in children]
|
||||||
|
|
||||||
|
def update_child(self, child, new_args):
|
||||||
|
args = getattr(child, '__discord_ui_model_kwargs__')
|
||||||
|
args |= new_args
|
||||||
|
|
||||||
|
def make_exports(self):
|
||||||
|
"""
|
||||||
|
Called post-instantiation to populate self._exports.
|
||||||
|
"""
|
||||||
|
return self.children
|
||||||
|
|
||||||
|
def refresh(self):
|
||||||
|
"""
|
||||||
|
Update widget components from current setting data, if applicable.
|
||||||
|
E.g. to update the default entry in a select list after a choice has been made,
|
||||||
|
or update button colours.
|
||||||
|
This does not trigger a discord ui update,
|
||||||
|
that is the responsibility of the interaction handler.
|
||||||
|
"""
|
||||||
|
pass
|
||||||
|
|
||||||
|
async def show(self, interaction: discord.Interaction, key: Any = None, override=False, **kwargs):
|
||||||
|
"""
|
||||||
|
Complete standard setting widget UI flow for this setting.
|
||||||
|
The SettingWidget components may be attached to other messages as needed,
|
||||||
|
and they may be triggered individually,
|
||||||
|
but this coroutine defines the standard interface.
|
||||||
|
Intended for use by any interaction which wants to "open the setting".
|
||||||
|
|
||||||
|
Extra keyword arguments are passed directly to the interaction reply (for e.g. ephemeral).
|
||||||
|
"""
|
||||||
|
if key is None:
|
||||||
|
# By default, only have one widget listener per interaction.
|
||||||
|
key = ('widget', interaction.id)
|
||||||
|
|
||||||
|
# If there is already a widget listening on this key, respect override
|
||||||
|
if self.setting.get_listener(key) and not override:
|
||||||
|
# Refuse to spawn another widget
|
||||||
|
return
|
||||||
|
|
||||||
|
async def update_callback(new_data):
|
||||||
|
self.setting.data = new_data
|
||||||
|
await interaction.edit_original_response(embed=self.setting.embed, view=self, **kwargs)
|
||||||
|
|
||||||
|
self.setting.register_callback(key)(update_callback)
|
||||||
|
await interaction.response.send_message(embed=self.setting.embed, view=self, **kwargs)
|
||||||
|
await self.wait()
|
||||||
|
try:
|
||||||
|
# Try and detach the view, since we aren't handling events anymore.
|
||||||
|
await interaction.edit_original_response(view=None)
|
||||||
|
except discord.HTTPException:
|
||||||
|
pass
|
||||||
|
self.setting.deregister_callback(key)
|
||||||
|
|
||||||
|
def attach(self, group_view: ui.View):
|
||||||
|
"""
|
||||||
|
Attach this setting widget to a view representing several settings.
|
||||||
|
"""
|
||||||
|
for item in self._exports:
|
||||||
|
group_view.add_item(item)
|
||||||
|
|
||||||
|
@button(style=ButtonStyle.secondary, label="Edit", row=4)
|
||||||
|
async def edit_button(self, interaction: discord.Interaction, button: ui.Button):
|
||||||
|
"""
|
||||||
|
Spawn a simple edit modal,
|
||||||
|
populated with `setting.input_field`.
|
||||||
|
"""
|
||||||
|
recover_context(self._context)
|
||||||
|
# Spawn the setting modal
|
||||||
|
await interaction.response.send_modal(self.modal)
|
||||||
|
|
||||||
|
@button(style=ButtonStyle.danger, label="Reset", row=4)
|
||||||
|
async def reset_button(self, interaction: discord.Interaction, button: Button):
|
||||||
|
recover_context(self._context)
|
||||||
|
await interaction.response.defer(thinking=True, ephemeral=True)
|
||||||
|
await self.setting.interactive_set(None, interaction)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def modal(self) -> Modal:
|
||||||
|
"""
|
||||||
|
Build a Modal dialogue for updating the setting.
|
||||||
|
Refreshes (and re-attaches) the input field each time this is called.
|
||||||
|
"""
|
||||||
|
if self._modal is not None:
|
||||||
|
self._modal.update_field(self.setting.input_field)
|
||||||
|
return self._modal
|
||||||
|
|
||||||
|
# TODO: Attach shared timeouts to the modal
|
||||||
|
self._modal = modal = SettingModal(
|
||||||
|
title=f"Edit {self.setting.display_name}",
|
||||||
|
)
|
||||||
|
modal.update_field(self.setting.input_field)
|
||||||
|
|
||||||
|
@modal.submit_callback()
|
||||||
|
async def edit_submit(interaction: discord.Interaction):
|
||||||
|
# TODO: Catch and handle UserInputError
|
||||||
|
await interaction.response.defer(thinking=True, ephemeral=True)
|
||||||
|
data = await self.setting._parse_string(self.setting.parent_id, modal.input_field.value)
|
||||||
|
await self.setting.interactive_set(data, interaction)
|
||||||
|
|
||||||
|
return modal
|
||||||
|
|
||||||
|
|
||||||
|
class InteractiveSetting(BaseSetting[ParentID, SettingData, SettingValue]):
|
||||||
|
__slots__ = ('_widget',)
|
||||||
|
|
||||||
|
# Configuration interface descriptions
|
||||||
|
_display_name: LazyStr # User readable name of the setting
|
||||||
|
_desc: LazyStr # User readable brief description of the setting
|
||||||
|
_long_desc: LazyStr # User readable long description of the setting
|
||||||
|
_accepts: LazyStr # User readable description of the acceptable values
|
||||||
|
_set_cmd: str = None
|
||||||
|
_notset_str: LazyStr = _p('setting|formatted|notset', "Not Set")
|
||||||
|
_virtual: bool = False # Whether the setting should be hidden from tables and dashboards
|
||||||
|
_required: bool = False
|
||||||
|
|
||||||
|
Widget = SettingWidget
|
||||||
|
|
||||||
|
# A list of callback coroutines to call when the setting updates
|
||||||
|
# This can be used globally to refresh state when the setting updates,
|
||||||
|
# Or locallly to e.g. refresh an active widget.
|
||||||
|
# The callbacks are called on write, so they may be bypassed by direct use of _writer!
|
||||||
|
_listeners_: Dict[Any, Callable[[Optional[SettingData]], Coroutine[Any, Any, None]]] = {}
|
||||||
|
|
||||||
|
# Optional client event to dispatch when theis setting has been written
|
||||||
|
# Event handlers should be of the form Callable[ParentID, SettingData]
|
||||||
|
_event: Optional[str] = None
|
||||||
|
|
||||||
|
# Interaction ward that should be validated via interaction_check
|
||||||
|
_write_ward: Optional[Callable[[discord.Interaction], Coroutine[Any, Any, bool]]] = None
|
||||||
|
|
||||||
|
def __init__(self, *args, **kwargs):
|
||||||
|
super().__init__(*args, **kwargs)
|
||||||
|
|
||||||
|
self._widget: Optional[SettingWidget] = None
|
||||||
|
|
||||||
|
@property
|
||||||
|
def long_desc(self):
|
||||||
|
t = ctx_translator.get().t
|
||||||
|
bot = ctx_bot.get()
|
||||||
|
return t(self._long_desc).format(
|
||||||
|
bot=bot,
|
||||||
|
cmds=bot.core.mention_cache
|
||||||
|
)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def display_name(self):
|
||||||
|
t = ctx_translator.get().t
|
||||||
|
return t(self._display_name)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def desc(self):
|
||||||
|
t = ctx_translator.get().t
|
||||||
|
return t(self._desc)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def accepts(self):
|
||||||
|
t = ctx_translator.get().t
|
||||||
|
return t(self._accepts)
|
||||||
|
|
||||||
|
async def write(self, **kwargs) -> None:
|
||||||
|
await super().write(**kwargs)
|
||||||
|
self.dispatch_update()
|
||||||
|
for listener in self._listeners_.values():
|
||||||
|
asyncio.create_task(listener(self.data))
|
||||||
|
|
||||||
|
def dispatch_update(self):
|
||||||
|
"""
|
||||||
|
Dispatch a client event along `self._event`, if set.
|
||||||
|
|
||||||
|
Override to modify the target event handler arguments.
|
||||||
|
By default, event handlers should be of the form:
|
||||||
|
Callable[[ParentID, SettingData], Coroutine[Any, Any, None]]
|
||||||
|
"""
|
||||||
|
if self._event is not None and (bot := ctx_bot.get()) is not None:
|
||||||
|
bot.dispatch(self._event, self.parent_id, self)
|
||||||
|
|
||||||
|
def get_listener(self, key):
|
||||||
|
return self._listeners_.get(key, None)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def register_callback(cls, name=None):
|
||||||
|
def wrapped(coro):
|
||||||
|
cls._listeners_[name or coro.__name__] = coro
|
||||||
|
return coro
|
||||||
|
return wrapped
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def deregister_callback(cls, name):
|
||||||
|
cls._listeners_.pop(name, None)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def update_message(self):
|
||||||
|
"""
|
||||||
|
Response message sent when the setting has successfully been updated.
|
||||||
|
Should generally be one line.
|
||||||
|
"""
|
||||||
|
if self.data is None:
|
||||||
|
return "Setting reset!"
|
||||||
|
else:
|
||||||
|
return f"Setting Updated! New value: {self.formatted}"
|
||||||
|
|
||||||
|
@property
|
||||||
|
def hover_desc(self):
|
||||||
|
"""
|
||||||
|
This no longer works since Discord changed the hover rules.
|
||||||
|
|
||||||
|
return '\n'.join((
|
||||||
|
self.display_name,
|
||||||
|
'=' * len(self.display_name),
|
||||||
|
self.desc,
|
||||||
|
f"\nAccepts: {self.accepts}"
|
||||||
|
))
|
||||||
|
"""
|
||||||
|
return self.desc
|
||||||
|
|
||||||
|
async def update_response(self, interaction: discord.Interaction, message: Optional[str] = None, **kwargs):
|
||||||
|
"""
|
||||||
|
Respond to an interaction which triggered a setting update.
|
||||||
|
Usually just wraps `update_message` in an embed and sends it back.
|
||||||
|
Passes any extra `kwargs` to the message creation method.
|
||||||
|
"""
|
||||||
|
embed = discord.Embed(
|
||||||
|
description=f"{str(conf.emojis.tick)} {message or self.update_message}",
|
||||||
|
colour=discord.Color.green()
|
||||||
|
)
|
||||||
|
if interaction.response.is_done():
|
||||||
|
await interaction.edit_original_response(embed=embed, **kwargs)
|
||||||
|
else:
|
||||||
|
await interaction.response.send_message(embed=embed, **kwargs)
|
||||||
|
|
||||||
|
async def interactive_set(self, new_data: Optional[SettingData], interaction: discord.Interaction, **kwargs):
|
||||||
|
self.data = new_data
|
||||||
|
await self.write()
|
||||||
|
await self.update_response(interaction, **kwargs)
|
||||||
|
|
||||||
|
async def format_in(self, bot, **kwargs):
|
||||||
|
"""
|
||||||
|
Formatted version of the setting given an asynchronous context with client.
|
||||||
|
"""
|
||||||
|
return self.formatted
|
||||||
|
|
||||||
|
@property
|
||||||
|
def embed_field(self):
|
||||||
|
"""
|
||||||
|
Returns a {name, value} pair for use in an Embed field.
|
||||||
|
"""
|
||||||
|
name = self.display_name
|
||||||
|
value = f"{self.long_desc}\n{self.desc_table}"
|
||||||
|
if len(value) > 1024:
|
||||||
|
t = ctx_translator.get().t
|
||||||
|
desc_table = '\n'.join(
|
||||||
|
tabulate(
|
||||||
|
*self._desc_table(
|
||||||
|
show_value=t(_p(
|
||||||
|
'setting|embed_field|too_long',
|
||||||
|
"Too long to display here!"
|
||||||
|
))
|
||||||
|
)
|
||||||
|
)
|
||||||
|
)
|
||||||
|
value = f"{self.long_desc}\n{desc_table}"
|
||||||
|
if len(value) > 1024:
|
||||||
|
# Forcibly trim
|
||||||
|
value = value[:1020] + '...'
|
||||||
|
return {'name': name, 'value': value}
|
||||||
|
|
||||||
|
@property
|
||||||
|
def set_str(self):
|
||||||
|
if self._set_cmd is not None:
|
||||||
|
bot = ctx_bot.get()
|
||||||
|
if bot:
|
||||||
|
return bot.core.mention_cmd(self._set_cmd)
|
||||||
|
else:
|
||||||
|
return f"`/{self._set_cmd}`"
|
||||||
|
|
||||||
|
@property
|
||||||
|
def notset_str(self):
|
||||||
|
t = ctx_translator.get().t
|
||||||
|
return t(self._notset_str)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def embed(self):
|
||||||
|
"""
|
||||||
|
Returns a full embed describing this setting.
|
||||||
|
"""
|
||||||
|
t = ctx_translator.get().t
|
||||||
|
embed = discord.Embed(
|
||||||
|
title=t(_p(
|
||||||
|
'setting|summary_embed|title',
|
||||||
|
"Configuration options for `{name}`"
|
||||||
|
)).format(name=self.display_name),
|
||||||
|
)
|
||||||
|
embed.description = "{}\n{}".format(self.long_desc.format(self=self), self.desc_table)
|
||||||
|
return embed
|
||||||
|
|
||||||
|
def _desc_table(self, show_value: Optional[str] = None) -> list[tuple[str, str]]:
|
||||||
|
t = ctx_translator.get().t
|
||||||
|
lines = []
|
||||||
|
|
||||||
|
# Currently line
|
||||||
|
lines.append((
|
||||||
|
t(_p('setting|summary_table|field:currently|key', "Currently")),
|
||||||
|
show_value or (self.formatted or self.notset_str)
|
||||||
|
))
|
||||||
|
|
||||||
|
# Default line
|
||||||
|
if (default := self.default) is not None:
|
||||||
|
lines.append((
|
||||||
|
t(_p('setting|summary_table|field:default|key', "By Default")),
|
||||||
|
self._format_data(self.parent_id, default) or 'None'
|
||||||
|
))
|
||||||
|
|
||||||
|
# Set using line
|
||||||
|
if (set_str := self.set_str) is not None:
|
||||||
|
lines.append((
|
||||||
|
t(_p('setting|summary_table|field:set|key', "Set Using")),
|
||||||
|
set_str
|
||||||
|
))
|
||||||
|
return lines
|
||||||
|
|
||||||
|
@property
|
||||||
|
def desc_table(self) -> str:
|
||||||
|
return '\n'.join(tabulate(*self._desc_table()))
|
||||||
|
|
||||||
|
@property
|
||||||
|
def input_field(self) -> TextInput:
|
||||||
|
"""
|
||||||
|
TextInput field used for string-based setting modification.
|
||||||
|
May be added to external modal for grouped setting editing.
|
||||||
|
This property is not persistent, and creates a new field each time.
|
||||||
|
"""
|
||||||
|
return TextInput(
|
||||||
|
label=self.display_name,
|
||||||
|
placeholder=self.accepts,
|
||||||
|
default=self.input_formatted[:4000] if self.input_formatted else None,
|
||||||
|
required=self._required
|
||||||
|
)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def widget(self):
|
||||||
|
"""
|
||||||
|
Returns the Discord UI View associated with the current setting.
|
||||||
|
"""
|
||||||
|
if self._widget is None:
|
||||||
|
self._widget = self.Widget(self)
|
||||||
|
return self._widget
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def set_widget(cls, WidgetCls):
|
||||||
|
"""
|
||||||
|
Convenience decorator to create the widget class for this setting.
|
||||||
|
"""
|
||||||
|
cls.Widget = WidgetCls
|
||||||
|
return WidgetCls
|
||||||
|
|
||||||
|
@property
|
||||||
|
def formatted(self):
|
||||||
|
"""
|
||||||
|
Default user-readable form of the setting.
|
||||||
|
Should be a short single line.
|
||||||
|
"""
|
||||||
|
return self._format_data(self.parent_id, self.data, **self.kwargs) or self.notset_str
|
||||||
|
|
||||||
|
@property
|
||||||
|
def input_formatted(self) -> str:
|
||||||
|
"""
|
||||||
|
Format the current value as a default value for an input field.
|
||||||
|
Returned string must be acceptable through parse_string.
|
||||||
|
Does not take into account defaults.
|
||||||
|
"""
|
||||||
|
if self._data is not None:
|
||||||
|
return str(self._data)
|
||||||
|
else:
|
||||||
|
return ""
|
||||||
|
|
||||||
|
@property
|
||||||
|
def summary(self):
|
||||||
|
"""
|
||||||
|
Formatted summary of the data.
|
||||||
|
May be implemented in `_format_data(..., summary=True, ...)` or overidden.
|
||||||
|
"""
|
||||||
|
return self._format_data(self.parent_id, self.data, summary=True, **self.kwargs)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
async def from_string(cls, parent_id, userstr: str, **kwargs):
|
||||||
|
"""
|
||||||
|
Return a setting instance initialised from a parsed user string.
|
||||||
|
"""
|
||||||
|
data = await cls._parse_string(parent_id, userstr, **kwargs)
|
||||||
|
return cls(parent_id, data, **kwargs)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
async def from_value(cls, parent_id, value, **kwargs):
|
||||||
|
await cls._check_value(parent_id, value, **kwargs)
|
||||||
|
data = cls._data_from_value(parent_id, value, **kwargs)
|
||||||
|
return cls(parent_id, data, **kwargs)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
async def _parse_string(cls, parent_id, string: str, **kwargs) -> Optional[SettingData]:
|
||||||
|
"""
|
||||||
|
Parse user provided string (usually from a TextInput) into raw setting data.
|
||||||
|
Must be overriden by the setting if the setting is user-configurable.
|
||||||
|
Returns None if the setting was unset.
|
||||||
|
"""
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _format_data(cls, parent_id, data, **kwargs):
|
||||||
|
"""
|
||||||
|
Convert raw setting data into a formatted user-readable string,
|
||||||
|
representing the current value.
|
||||||
|
"""
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
async def _check_value(cls, parent_id, value, **kwargs):
|
||||||
|
"""
|
||||||
|
Check the provided value is valid.
|
||||||
|
|
||||||
|
Many setting update methods now provide Discord objects instead of raw data or user strings.
|
||||||
|
This method may be used for value-checking such a value.
|
||||||
|
|
||||||
|
Raises UserInputError if the value fails validation.
|
||||||
|
"""
|
||||||
|
pass
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
async def interaction_check(cls, parent_id, interaction: discord.Interaction, **kwargs):
|
||||||
|
if cls._write_ward is not None and not await cls._write_ward(interaction):
|
||||||
|
# TODO: Combine the check system so we can do customised errors here
|
||||||
|
t = ctx_translator.get().t
|
||||||
|
raise UserInputError(t(_p(
|
||||||
|
'setting|interaction_check|error',
|
||||||
|
"You do not have sufficient permissions to do this!"
|
||||||
|
)))
|
||||||
|
|
||||||
|
|
||||||
|
"""
|
||||||
|
command callback for set command?
|
||||||
|
autocomplete for set command?
|
||||||
|
|
||||||
|
Might be better in a ConfigSetting subclass.
|
||||||
|
But also mix into the base setting types.
|
||||||
|
"""
|
||||||
+304
-106
@@ -1,5 +1,6 @@
|
|||||||
from io import StringIO
|
from io import StringIO
|
||||||
from typing import NamedTuple, Optional, Sequence, Union, overload, List, Any
|
from typing import NamedTuple, Optional, Sequence, Union, overload, List, Any
|
||||||
|
import asyncio
|
||||||
import collections
|
import collections
|
||||||
import datetime
|
import datetime
|
||||||
import datetime as dt
|
import datetime as dt
|
||||||
@@ -11,18 +12,24 @@ from contextvars import Context
|
|||||||
|
|
||||||
import discord
|
import discord
|
||||||
from discord.partial_emoji import _EmojiTag
|
from discord.partial_emoji import _EmojiTag
|
||||||
from discord import Embed, File, GuildSticker, StickerItem, AllowedMentions, Message, MessageReference, PartialMessage
|
from discord import (
|
||||||
|
Embed,
|
||||||
|
File,
|
||||||
|
GuildSticker,
|
||||||
|
StickerItem,
|
||||||
|
AllowedMentions,
|
||||||
|
Message,
|
||||||
|
MessageReference,
|
||||||
|
PartialMessage,
|
||||||
|
)
|
||||||
from discord.ui import View
|
from discord.ui import View
|
||||||
|
|
||||||
from meta.errors import UserInputError
|
from meta.errors import UserInputError
|
||||||
|
|
||||||
|
|
||||||
multiselect_regex = re.compile(
|
multiselect_regex = re.compile(r"^([0-9, -]+)$", re.DOTALL | re.IGNORECASE | re.VERBOSE)
|
||||||
r"^([0-9, -]+)$",
|
tick = "✅"
|
||||||
re.DOTALL | re.IGNORECASE | re.VERBOSE
|
cross = "❌"
|
||||||
)
|
|
||||||
tick = '✅'
|
|
||||||
cross = '❌'
|
|
||||||
|
|
||||||
MISSING = object()
|
MISSING = object()
|
||||||
|
|
||||||
@@ -31,6 +38,7 @@ class MessageArgs:
|
|||||||
"""
|
"""
|
||||||
Utility class for storing message creation and editing arguments.
|
Utility class for storing message creation and editing arguments.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
# TODO: Overrides for mutually exclusive arguments, see Messageable.send
|
# TODO: Overrides for mutually exclusive arguments, see Messageable.send
|
||||||
|
|
||||||
@overload
|
@overload
|
||||||
@@ -49,8 +57,7 @@ class MessageArgs:
|
|||||||
mention_author: bool = ...,
|
mention_author: bool = ...,
|
||||||
view: View = ...,
|
view: View = ...,
|
||||||
suppress_embeds: bool = ...,
|
suppress_embeds: bool = ...,
|
||||||
) -> None:
|
) -> None: ...
|
||||||
...
|
|
||||||
|
|
||||||
@overload
|
@overload
|
||||||
def __init__(
|
def __init__(
|
||||||
@@ -68,8 +75,7 @@ class MessageArgs:
|
|||||||
mention_author: bool = ...,
|
mention_author: bool = ...,
|
||||||
view: View = ...,
|
view: View = ...,
|
||||||
suppress_embeds: bool = ...,
|
suppress_embeds: bool = ...,
|
||||||
) -> None:
|
) -> None: ...
|
||||||
...
|
|
||||||
|
|
||||||
@overload
|
@overload
|
||||||
def __init__(
|
def __init__(
|
||||||
@@ -87,8 +93,7 @@ class MessageArgs:
|
|||||||
mention_author: bool = ...,
|
mention_author: bool = ...,
|
||||||
view: View = ...,
|
view: View = ...,
|
||||||
suppress_embeds: bool = ...,
|
suppress_embeds: bool = ...,
|
||||||
) -> None:
|
) -> None: ...
|
||||||
...
|
|
||||||
|
|
||||||
@overload
|
@overload
|
||||||
def __init__(
|
def __init__(
|
||||||
@@ -106,17 +111,16 @@ class MessageArgs:
|
|||||||
mention_author: bool = ...,
|
mention_author: bool = ...,
|
||||||
view: View = ...,
|
view: View = ...,
|
||||||
suppress_embeds: bool = ...,
|
suppress_embeds: bool = ...,
|
||||||
) -> None:
|
) -> None: ...
|
||||||
...
|
|
||||||
|
|
||||||
def __init__(self, **kwargs):
|
def __init__(self, **kwargs):
|
||||||
self.kwargs = kwargs
|
self.kwargs = kwargs
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def send_args(self) -> dict:
|
def send_args(self) -> dict:
|
||||||
if self.kwargs.get('view', MISSING) is None:
|
if self.kwargs.get("view", MISSING) is None:
|
||||||
kwargs = self.kwargs.copy()
|
kwargs = self.kwargs.copy()
|
||||||
kwargs.pop('view')
|
kwargs.pop("view")
|
||||||
else:
|
else:
|
||||||
kwargs = self.kwargs
|
kwargs = self.kwargs
|
||||||
|
|
||||||
@@ -126,20 +130,25 @@ class MessageArgs:
|
|||||||
def edit_args(self) -> dict:
|
def edit_args(self) -> dict:
|
||||||
args = {}
|
args = {}
|
||||||
kept = (
|
kept = (
|
||||||
'content', 'embed', 'embeds', 'delete_after', 'allowed_mentions', 'view'
|
"content",
|
||||||
|
"embed",
|
||||||
|
"embeds",
|
||||||
|
"delete_after",
|
||||||
|
"allowed_mentions",
|
||||||
|
"view",
|
||||||
)
|
)
|
||||||
for k in kept:
|
for k in kept:
|
||||||
if k in self.kwargs:
|
if k in self.kwargs:
|
||||||
args[k] = self.kwargs[k]
|
args[k] = self.kwargs[k]
|
||||||
|
|
||||||
if 'file' in self.kwargs:
|
if "file" in self.kwargs:
|
||||||
args['attachments'] = [self.kwargs['file']]
|
args["attachments"] = [self.kwargs["file"]]
|
||||||
|
|
||||||
if 'files' in self.kwargs:
|
if "files" in self.kwargs:
|
||||||
args['attachments'] = self.kwargs['files']
|
args["attachments"] = self.kwargs["files"]
|
||||||
|
|
||||||
if 'suppress_embeds' in self.kwargs:
|
if "suppress_embeds" in self.kwargs:
|
||||||
args['suppress'] = self.kwargs['suppress_embeds']
|
args["suppress"] = self.kwargs["suppress_embeds"]
|
||||||
|
|
||||||
return args
|
return args
|
||||||
|
|
||||||
@@ -148,9 +157,9 @@ def tabulate(
|
|||||||
*fields: tuple[str, str],
|
*fields: tuple[str, str],
|
||||||
row_format: str = "`{invis}{key:<{pad}}{colon}`\t{value}",
|
row_format: str = "`{invis}{key:<{pad}}{colon}`\t{value}",
|
||||||
sub_format: str = "`{invis:<{pad}}{colon}`\t{value}",
|
sub_format: str = "`{invis:<{pad}}{colon}`\t{value}",
|
||||||
colon: str = ':',
|
colon: str = ":",
|
||||||
invis: str = "",
|
invis: str = "",
|
||||||
**args
|
**args,
|
||||||
) -> list[str]:
|
) -> list[str]:
|
||||||
"""
|
"""
|
||||||
Turns a list of (property, value) pairs into
|
Turns a list of (property, value) pairs into
|
||||||
@@ -181,7 +190,7 @@ def tabulate(
|
|||||||
for field in fields:
|
for field in fields:
|
||||||
key = field[0]
|
key = field[0]
|
||||||
value = field[1]
|
value = field[1]
|
||||||
lines = value.split('\r\n')
|
lines = value.split("\r\n")
|
||||||
|
|
||||||
row_line = row_format.format(
|
row_line = row_format.format(
|
||||||
invis=invis,
|
invis=invis,
|
||||||
@@ -190,7 +199,7 @@ def tabulate(
|
|||||||
colon=colon,
|
colon=colon,
|
||||||
value=lines[0],
|
value=lines[0],
|
||||||
field=field,
|
field=field,
|
||||||
**args
|
**args,
|
||||||
)
|
)
|
||||||
if len(lines) > 1:
|
if len(lines) > 1:
|
||||||
row_lines = [row_line]
|
row_lines = [row_line]
|
||||||
@@ -200,15 +209,17 @@ def tabulate(
|
|||||||
pad=max_len + len(colon),
|
pad=max_len + len(colon),
|
||||||
colon=colon,
|
colon=colon,
|
||||||
value=line,
|
value=line,
|
||||||
**args
|
**args,
|
||||||
)
|
)
|
||||||
row_lines.append(sub_line)
|
row_lines.append(sub_line)
|
||||||
row_line = '\n'.join(row_lines)
|
row_line = "\n".join(row_lines)
|
||||||
rows.append(row_line)
|
rows.append(row_line)
|
||||||
return rows
|
return rows
|
||||||
|
|
||||||
|
|
||||||
def paginate_list(item_list: list[str], block_length=20, style="markdown", title=None) -> list[str]:
|
def paginate_list(
|
||||||
|
item_list: list[str], block_length=20, style="markdown", title=None
|
||||||
|
) -> list[str]:
|
||||||
"""
|
"""
|
||||||
Create pretty codeblock pages from a list of strings.
|
Create pretty codeblock pages from a list of strings.
|
||||||
|
|
||||||
@@ -229,8 +240,13 @@ def paginate_list(item_list: list[str], block_length=20, style="markdown", title
|
|||||||
List of pages, each formatted into a codeblock,
|
List of pages, each formatted into a codeblock,
|
||||||
and containing at most `block_length` of the provided strings.
|
and containing at most `block_length` of the provided strings.
|
||||||
"""
|
"""
|
||||||
lines = ["{0:<5}{1:<5}".format("{}.".format(i + 1), str(line)) for i, line in enumerate(item_list)]
|
lines = [
|
||||||
page_blocks = [lines[i:i + block_length] for i in range(0, len(lines), block_length)]
|
"{0:<5}{1:<5}".format("{}.".format(i + 1), str(line))
|
||||||
|
for i, line in enumerate(item_list)
|
||||||
|
]
|
||||||
|
page_blocks = [
|
||||||
|
lines[i : i + block_length] for i in range(0, len(lines), block_length)
|
||||||
|
]
|
||||||
pages = []
|
pages = []
|
||||||
for i, block in enumerate(page_blocks):
|
for i, block in enumerate(page_blocks):
|
||||||
pagenum = "Page {}/{}".format(i + 1, len(page_blocks))
|
pagenum = "Page {}/{}".format(i + 1, len(page_blocks))
|
||||||
@@ -239,12 +255,18 @@ def paginate_list(item_list: list[str], block_length=20, style="markdown", title
|
|||||||
else:
|
else:
|
||||||
header = pagenum
|
header = pagenum
|
||||||
header_line = "=" * len(header)
|
header_line = "=" * len(header)
|
||||||
full_header = "{}\n{}\n".format(header, header_line) if len(page_blocks) > 1 or title else ""
|
full_header = (
|
||||||
|
"{}\n{}\n".format(header, header_line)
|
||||||
|
if len(page_blocks) > 1 or title
|
||||||
|
else ""
|
||||||
|
)
|
||||||
pages.append("```{}\n{}{}```".format(style, full_header, "\n".join(block)))
|
pages.append("```{}\n{}{}```".format(style, full_header, "\n".join(block)))
|
||||||
return pages
|
return pages
|
||||||
|
|
||||||
|
|
||||||
def split_text(text: str, blocksize=2000, code=True, syntax="", maxheight=50) -> list[str]:
|
def split_text(
|
||||||
|
text: str, blocksize=2000, code=True, syntax="", maxheight=50
|
||||||
|
) -> list[str]:
|
||||||
"""
|
"""
|
||||||
Break the text into blocks of maximum length blocksize
|
Break the text into blocks of maximum length blocksize
|
||||||
If possible, break across nearby newlines. Otherwise just break at blocksize chars
|
If possible, break across nearby newlines. Otherwise just break at blocksize chars
|
||||||
@@ -277,10 +299,10 @@ def split_text(text: str, blocksize=2000, code=True, syntax="", maxheight=50) ->
|
|||||||
if len(text) <= blocksize:
|
if len(text) <= blocksize:
|
||||||
blocks.append(text)
|
blocks.append(text)
|
||||||
break
|
break
|
||||||
text = text.strip('\n')
|
text = text.strip("\n")
|
||||||
|
|
||||||
# Find the last newline in the prototype block
|
# Find the last newline in the prototype block
|
||||||
split_on = text[0:blocksize].rfind('\n')
|
split_on = text[0:blocksize].rfind("\n")
|
||||||
split_on = blocksize if split_on < blocksize // 5 else split_on
|
split_on = blocksize if split_on < blocksize // 5 else split_on
|
||||||
|
|
||||||
# Add the block and truncate the text
|
# Add the block and truncate the text
|
||||||
@@ -313,15 +335,17 @@ def strfdelta(delta: datetime.timedelta, sec=False, minutes=True, short=False) -
|
|||||||
A string containing a time from the datetime.timedelta object, in a readable format.
|
A string containing a time from the datetime.timedelta object, in a readable format.
|
||||||
Time units will be abbreviated if short was set to True.
|
Time units will be abbreviated if short was set to True.
|
||||||
"""
|
"""
|
||||||
output = [[delta.days, 'd' if short else ' day'],
|
output = [
|
||||||
[delta.seconds // 3600, 'h' if short else ' hour']]
|
[delta.days, "d" if short else " day"],
|
||||||
|
[delta.seconds // 3600, "h" if short else " hour"],
|
||||||
|
]
|
||||||
if minutes:
|
if minutes:
|
||||||
output.append([delta.seconds // 60 % 60, 'm' if short else ' minute'])
|
output.append([delta.seconds // 60 % 60, "m" if short else " minute"])
|
||||||
if sec:
|
if sec:
|
||||||
output.append([delta.seconds % 60, 's' if short else ' second'])
|
output.append([delta.seconds % 60, "s" if short else " second"])
|
||||||
for i in range(len(output)):
|
for i in range(len(output)):
|
||||||
if output[i][0] != 1 and not short:
|
if output[i][0] != 1 and not short:
|
||||||
output[i][1] += 's' # type: ignore
|
output[i][1] += "s" # type: ignore
|
||||||
reply_msg = []
|
reply_msg = []
|
||||||
if output[0][0] != 0:
|
if output[0][0] != 0:
|
||||||
reply_msg.append("{}{} ".format(output[0][0], output[0][1]))
|
reply_msg.append("{}{} ".format(output[0][0], output[0][1]))
|
||||||
@@ -347,12 +371,14 @@ def _parse_dur(time_str: str) -> int:
|
|||||||
Returns: int
|
Returns: int
|
||||||
The number of seconds the duration represents.
|
The number of seconds the duration represents.
|
||||||
"""
|
"""
|
||||||
funcs = {'d': lambda x: x * 24 * 60 * 60,
|
funcs = {
|
||||||
'h': lambda x: x * 60 * 60,
|
"d": lambda x: x * 24 * 60 * 60,
|
||||||
'm': lambda x: x * 60,
|
"h": lambda x: x * 60 * 60,
|
||||||
's': lambda x: x}
|
"m": lambda x: x * 60,
|
||||||
|
"s": lambda x: x,
|
||||||
|
}
|
||||||
time_str = time_str.strip(" ,")
|
time_str = time_str.strip(" ,")
|
||||||
found = re.findall(r'(\d+)\s?(\w+?)', time_str)
|
found = re.findall(r"(\d+)\s?(\w+?)", time_str)
|
||||||
seconds = 0
|
seconds = 0
|
||||||
for bit in found:
|
for bit in found:
|
||||||
if bit[1] in funcs:
|
if bit[1] in funcs:
|
||||||
@@ -373,25 +399,27 @@ def strfdur(duration: int, short=True, show_days=False) -> str:
|
|||||||
|
|
||||||
parts = []
|
parts = []
|
||||||
if days:
|
if days:
|
||||||
unit = 'd' if short else (' days' if days != 1 else ' day')
|
unit = "d" if short else (" days" if days != 1 else " day")
|
||||||
parts.append('{}{}'.format(days, unit))
|
parts.append("{}{}".format(days, unit))
|
||||||
if hours:
|
if hours:
|
||||||
unit = 'h' if short else (' hours' if hours != 1 else ' hour')
|
unit = "h" if short else (" hours" if hours != 1 else " hour")
|
||||||
parts.append('{}{}'.format(hours, unit))
|
parts.append("{}{}".format(hours, unit))
|
||||||
if minutes:
|
if minutes:
|
||||||
unit = 'm' if short else (' minutes' if minutes != 1 else ' minute')
|
unit = "m" if short else (" minutes" if minutes != 1 else " minute")
|
||||||
parts.append('{}{}'.format(minutes, unit))
|
parts.append("{}{}".format(minutes, unit))
|
||||||
if seconds or duration == 0:
|
if seconds or duration == 0:
|
||||||
unit = 's' if short else (' seconds' if seconds != 1 else ' second')
|
unit = "s" if short else (" seconds" if seconds != 1 else " second")
|
||||||
parts.append('{}{}'.format(seconds, unit))
|
parts.append("{}{}".format(seconds, unit))
|
||||||
|
|
||||||
if short:
|
if short:
|
||||||
return ' '.join(parts)
|
return " ".join(parts)
|
||||||
else:
|
else:
|
||||||
return ', '.join(parts)
|
return ", ".join(parts)
|
||||||
|
|
||||||
|
|
||||||
def substitute_ranges(ranges_str: str, max_match=20, max_range=1000, separator=',') -> str:
|
def substitute_ranges(
|
||||||
|
ranges_str: str, max_match=20, max_range=1000, separator=","
|
||||||
|
) -> str:
|
||||||
"""
|
"""
|
||||||
Substitutes a user provided list of numbers and ranges,
|
Substitutes a user provided list of numbers and ranges,
|
||||||
and replaces the ranges by the corresponding list of numbers.
|
and replaces the ranges by the corresponding list of numbers.
|
||||||
@@ -407,6 +435,7 @@ def substitute_ranges(ranges_str: str, max_match=20, max_range=1000, separator='
|
|||||||
The maximum length of range to replace.
|
The maximum length of range to replace.
|
||||||
Attempting to replace a range longer than this will raise a `ValueError`.
|
Attempting to replace a range longer than this will raise a `ValueError`.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def _repl(match):
|
def _repl(match):
|
||||||
n1 = int(match.group(1))
|
n1 = int(match.group(1))
|
||||||
n2 = int(match.group(2))
|
n2 = int(match.group(2))
|
||||||
@@ -415,16 +444,18 @@ def substitute_ranges(ranges_str: str, max_match=20, max_range=1000, separator='
|
|||||||
raise ValueError("Provided range is too large!")
|
raise ValueError("Provided range is too large!")
|
||||||
return separator.join(str(i) for i in range(n1, n2 + 1))
|
return separator.join(str(i) for i in range(n1, n2 + 1))
|
||||||
|
|
||||||
return re.sub(r'(\d+)\s*-\s*(\d+)', _repl, ranges_str, max_match)
|
return re.sub(r"(\d+)\s*-\s*(\d+)", _repl, ranges_str, max_match)
|
||||||
|
|
||||||
|
|
||||||
def parse_ranges(ranges_str: str, ignore_errors=False, separator=',', **kwargs) -> list[int]:
|
def parse_ranges(
|
||||||
|
ranges_str: str, ignore_errors=False, separator=",", **kwargs
|
||||||
|
) -> list[int]:
|
||||||
"""
|
"""
|
||||||
Parses a user provided range string into a list of numbers.
|
Parses a user provided range string into a list of numbers.
|
||||||
Extra keyword arguments are transparently passed to the underlying parser `substitute_ranges`.
|
Extra keyword arguments are transparently passed to the underlying parser `substitute_ranges`.
|
||||||
"""
|
"""
|
||||||
substituted = substitute_ranges(ranges_str, separator=separator, **kwargs)
|
substituted = substitute_ranges(ranges_str, separator=separator, **kwargs)
|
||||||
_numbers = (item.strip() for item in substituted.split(','))
|
_numbers = (item.strip() for item in substituted.split(","))
|
||||||
numbers = [item for item in _numbers if item]
|
numbers = [item for item in _numbers if item]
|
||||||
integers = [int(item) for item in numbers if item.isdigit()]
|
integers = [int(item) for item in numbers if item.isdigit()]
|
||||||
|
|
||||||
@@ -438,7 +469,9 @@ def parse_ranges(ranges_str: str, ignore_errors=False, separator=',', **kwargs)
|
|||||||
return integers
|
return integers
|
||||||
|
|
||||||
|
|
||||||
def msg_string(msg: discord.Message, mask_link=False, line_break=False, tz=None, clean=True) -> str:
|
def msg_string(
|
||||||
|
msg: discord.Message, mask_link=False, line_break=False, tz=None, clean=True
|
||||||
|
) -> str:
|
||||||
"""
|
"""
|
||||||
Format a message into a string with various information, such as:
|
Format a message into a string with various information, such as:
|
||||||
the timestamp of the message, author, message content, and attachments.
|
the timestamp of the message, author, message content, and attachments.
|
||||||
@@ -462,20 +495,26 @@ def msg_string(msg: discord.Message, mask_link=False, line_break=False, tz=None,
|
|||||||
"""
|
"""
|
||||||
timestr = "%I:%M %p, %d/%m/%Y"
|
timestr = "%I:%M %p, %d/%m/%Y"
|
||||||
if tz:
|
if tz:
|
||||||
time = iso8601.parse_date(msg.created_at.isoformat()).astimezone(tz).strftime(timestr)
|
time = (
|
||||||
|
iso8601.parse_date(msg.created_at.isoformat())
|
||||||
|
.astimezone(tz)
|
||||||
|
.strftime(timestr)
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
time = msg.created_at.strftime(timestr)
|
time = msg.created_at.strftime(timestr)
|
||||||
user = str(msg.author)
|
user = str(msg.author)
|
||||||
attach_list = [attach.proxy_url for attach in msg.attachments if attach.proxy_url]
|
attach_list = [attach.proxy_url for attach in msg.attachments if attach.proxy_url]
|
||||||
if mask_link:
|
if mask_link:
|
||||||
attach_list = ["[Link]({})".format(url) for url in attach_list]
|
attach_list = ["[Link]({})".format(url) for url in attach_list]
|
||||||
attachments = "\nAttachments: {}".format(", ".join(attach_list)) if attach_list else ""
|
attachments = (
|
||||||
|
"\nAttachments: {}".format(", ".join(attach_list)) if attach_list else ""
|
||||||
|
)
|
||||||
return "`[{time}]` **{user}:** {line_break}{message} {attachments}".format(
|
return "`[{time}]` **{user}:** {line_break}{message} {attachments}".format(
|
||||||
time=time,
|
time=time,
|
||||||
user=user,
|
user=user,
|
||||||
line_break="\n" if line_break else "",
|
line_break="\n" if line_break else "",
|
||||||
message=msg.clean_content if clean else msg.content,
|
message=msg.clean_content if clean else msg.content,
|
||||||
attachments=attachments
|
attachments=attachments,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -491,21 +530,23 @@ def convdatestring(datestring: str) -> datetime.timedelta:
|
|||||||
Returns: datetime.timedelta
|
Returns: datetime.timedelta
|
||||||
A datetime.timedelta object formed from the string provided.
|
A datetime.timedelta object formed from the string provided.
|
||||||
"""
|
"""
|
||||||
datestring = datestring.strip(' ,')
|
datestring = datestring.strip(" ,")
|
||||||
datearray = []
|
datearray = []
|
||||||
funcs = {'d': lambda x: x * 24 * 60 * 60,
|
funcs = {
|
||||||
'h': lambda x: x * 60 * 60,
|
"d": lambda x: x * 24 * 60 * 60,
|
||||||
'm': lambda x: x * 60,
|
"h": lambda x: x * 60 * 60,
|
||||||
's': lambda x: x}
|
"m": lambda x: x * 60,
|
||||||
currentnumber = ''
|
"s": lambda x: x,
|
||||||
|
}
|
||||||
|
currentnumber = ""
|
||||||
for char in datestring:
|
for char in datestring:
|
||||||
if char.isdigit():
|
if char.isdigit():
|
||||||
currentnumber += char
|
currentnumber += char
|
||||||
else:
|
else:
|
||||||
if currentnumber == '':
|
if currentnumber == "":
|
||||||
continue
|
continue
|
||||||
datearray.append((int(currentnumber), char))
|
datearray.append((int(currentnumber), char))
|
||||||
currentnumber = ''
|
currentnumber = ""
|
||||||
seconds = 0
|
seconds = 0
|
||||||
if currentnumber:
|
if currentnumber:
|
||||||
seconds += int(currentnumber)
|
seconds += int(currentnumber)
|
||||||
@@ -520,6 +561,7 @@ class _rawChannel(discord.abc.Messageable):
|
|||||||
Raw messageable class representing an arbitrary channel,
|
Raw messageable class representing an arbitrary channel,
|
||||||
not necessarially seen by the gateway.
|
not necessarially seen by the gateway.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, state, id):
|
def __init__(self, state, id):
|
||||||
self._state = state
|
self._state = state
|
||||||
self.id = id
|
self.id = id
|
||||||
@@ -586,8 +628,12 @@ def join_list(string: list[str], nfs=False) -> str:
|
|||||||
"""
|
"""
|
||||||
# TODO: Probably not useful with localisation
|
# TODO: Probably not useful with localisation
|
||||||
if len(string) > 1:
|
if len(string) > 1:
|
||||||
return "{}{} and {}{}".format((", ").join(string[:-1]),
|
return "{}{} and {}{}".format(
|
||||||
"," if len(string) > 2 else "", string[-1], "" if nfs else ".")
|
(", ").join(string[:-1]),
|
||||||
|
"," if len(string) > 2 else "",
|
||||||
|
string[-1],
|
||||||
|
"" if nfs else ".",
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
return "{}{}".format("".join(string), "" if nfs else ".")
|
return "{}{}".format("".join(string), "" if nfs else ".")
|
||||||
|
|
||||||
@@ -603,10 +649,8 @@ def jumpto(guildid: int, channeldid: int, messageid: int) -> str:
|
|||||||
"""
|
"""
|
||||||
Build a jump link for a message given its location.
|
Build a jump link for a message given its location.
|
||||||
"""
|
"""
|
||||||
return 'https://discord.com/channels/{}/{}/{}'.format(
|
return "https://discord.com/channels/{}/{}/{}".format(
|
||||||
guildid,
|
guildid, channeldid, messageid
|
||||||
channeldid,
|
|
||||||
messageid
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -621,7 +665,7 @@ def multiple_replace(string: str, rep_dict: dict[str, str]) -> str:
|
|||||||
if rep_dict:
|
if rep_dict:
|
||||||
pattern = re.compile(
|
pattern = re.compile(
|
||||||
"|".join([re.escape(k) for k in sorted(rep_dict, key=len, reverse=True)]),
|
"|".join([re.escape(k) for k in sorted(rep_dict, key=len, reverse=True)]),
|
||||||
flags=re.DOTALL
|
flags=re.DOTALL,
|
||||||
)
|
)
|
||||||
return pattern.sub(lambda x: str(rep_dict[x.group(0)]), string)
|
return pattern.sub(lambda x: str(rep_dict[x.group(0)]), string)
|
||||||
else:
|
else:
|
||||||
@@ -644,12 +688,16 @@ def parse_ids(idstr: str) -> List[int]:
|
|||||||
from meta.errors import UserInputError
|
from meta.errors import UserInputError
|
||||||
|
|
||||||
# Extract ids from string
|
# Extract ids from string
|
||||||
splititer = (split.strip('<@!#&>, ') for split in idstr.split(','))
|
splititer = (split.strip("<@!#&>, ") for split in idstr.split(","))
|
||||||
splits = [split for split in splititer if split]
|
splits = [split for split in splititer if split]
|
||||||
|
|
||||||
# Check they are integers
|
# Check they are integers
|
||||||
if (not_id := next((split for split in splits if not split.isdigit()), None)) is not None:
|
if (
|
||||||
raise UserInputError("Could not extract an id from `$item`!", {'orig': idstr, 'item': not_id})
|
not_id := next((split for split in splits if not split.isdigit()), None)
|
||||||
|
) is not None:
|
||||||
|
raise UserInputError(
|
||||||
|
"Could not extract an id from `$item`!", {"orig": idstr, "item": not_id}
|
||||||
|
)
|
||||||
|
|
||||||
# Cast to integer and return
|
# Cast to integer and return
|
||||||
return list(map(int, splits))
|
return list(map(int, splits))
|
||||||
@@ -657,9 +705,7 @@ def parse_ids(idstr: str) -> List[int]:
|
|||||||
|
|
||||||
def error_embed(error, **kwargs) -> discord.Embed:
|
def error_embed(error, **kwargs) -> discord.Embed:
|
||||||
embed = discord.Embed(
|
embed = discord.Embed(
|
||||||
colour=discord.Colour.brand_red(),
|
colour=discord.Colour.brand_red(), description=error, timestamp=utc_now()
|
||||||
description=error,
|
|
||||||
timestamp=utc_now()
|
|
||||||
)
|
)
|
||||||
return embed
|
return embed
|
||||||
|
|
||||||
@@ -676,6 +722,7 @@ class Timezoned:
|
|||||||
|
|
||||||
Provides several useful localised properties.
|
Provides several useful localised properties.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
__slots__ = ()
|
__slots__ = ()
|
||||||
|
|
||||||
@property
|
@property
|
||||||
@@ -727,8 +774,10 @@ def replace_multiple(format_string, mapping):
|
|||||||
raise ValueError("Empty mapping passed.")
|
raise ValueError("Empty mapping passed.")
|
||||||
|
|
||||||
keys = list(mapping.keys())
|
keys = list(mapping.keys())
|
||||||
pattern = '|'.join(f"({key})" for key in keys)
|
pattern = "|".join(f"({key})" for key in keys)
|
||||||
string = re.sub(pattern, lambda match: str(mapping[keys[match.lastindex - 1]]), format_string)
|
string = re.sub(
|
||||||
|
pattern, lambda match: str(mapping[keys[match.lastindex - 1]]), format_string
|
||||||
|
)
|
||||||
return string
|
return string
|
||||||
|
|
||||||
|
|
||||||
@@ -748,6 +797,7 @@ def emojikey(emoji: discord.Emoji | discord.PartialEmoji | str):
|
|||||||
|
|
||||||
return key
|
return key
|
||||||
|
|
||||||
|
|
||||||
def recurse_map(func, obj, loc=[]):
|
def recurse_map(func, obj, loc=[]):
|
||||||
if isinstance(obj, dict):
|
if isinstance(obj, dict):
|
||||||
for k, v in obj.items():
|
for k, v in obj.items():
|
||||||
@@ -761,7 +811,8 @@ def recurse_map(func, obj, loc=[]):
|
|||||||
loc.pop()
|
loc.pop()
|
||||||
else:
|
else:
|
||||||
obj = func(loc, obj)
|
obj = func(loc, obj)
|
||||||
return obj
|
return obj
|
||||||
|
|
||||||
|
|
||||||
async def check_dm(user: discord.User | discord.Member) -> bool:
|
async def check_dm(user: discord.User | discord.Member) -> bool:
|
||||||
"""
|
"""
|
||||||
@@ -774,7 +825,7 @@ async def check_dm(user: discord.User | discord.Member) -> bool:
|
|||||||
(i.e. during a user instigated interaction).
|
(i.e. during a user instigated interaction).
|
||||||
"""
|
"""
|
||||||
try:
|
try:
|
||||||
await user.send('')
|
await user.send("")
|
||||||
except discord.Forbidden:
|
except discord.Forbidden:
|
||||||
return False
|
return False
|
||||||
except discord.HTTPException:
|
except discord.HTTPException:
|
||||||
@@ -783,38 +834,36 @@ async def check_dm(user: discord.User | discord.Member) -> bool:
|
|||||||
|
|
||||||
async def command_lengths(tree) -> dict[str, int]:
|
async def command_lengths(tree) -> dict[str, int]:
|
||||||
cmds = tree.get_commands()
|
cmds = tree.get_commands()
|
||||||
payloads = [
|
payloads = [await cmd.get_translated_payload(tree.translator) for cmd in cmds]
|
||||||
await cmd.get_translated_payload(tree.translator)
|
|
||||||
for cmd in cmds
|
|
||||||
]
|
|
||||||
lens = {}
|
lens = {}
|
||||||
for command in payloads:
|
for command in payloads:
|
||||||
name = command['name']
|
name = command["name"]
|
||||||
crumbs = {}
|
crumbs = {}
|
||||||
cmd_len = lens[name] = _recurse_length(command, crumbs, (name,))
|
cmd_len = lens[name] = _recurse_length(command, crumbs, (name,))
|
||||||
if name == 'configure' or cmd_len > 4000:
|
if name == "configure" or cmd_len > 4000:
|
||||||
print(f"'{name}' over 4000. Breadcrumb Trail follows:")
|
print(f"'{name}' over 4000. Breadcrumb Trail follows:")
|
||||||
lines = []
|
lines = []
|
||||||
for loc, val in crumbs.items():
|
for loc, val in crumbs.items():
|
||||||
locstr = '.'.join(loc)
|
locstr = ".".join(loc)
|
||||||
lines.append(f"{locstr}: {val}")
|
lines.append(f"{locstr}: {val}")
|
||||||
print('\n'.join(lines))
|
print("\n".join(lines))
|
||||||
print(json.dumps(command, indent=2))
|
print(json.dumps(command, indent=2))
|
||||||
return lens
|
return lens
|
||||||
|
|
||||||
|
|
||||||
def _recurse_length(payload, breadcrumbs={}, header=()) -> int:
|
def _recurse_length(payload, breadcrumbs={}, header=()) -> int:
|
||||||
total = 0
|
total = 0
|
||||||
total_header = (*header, '')
|
total_header = (*header, "")
|
||||||
breadcrumbs[total_header] = 0
|
breadcrumbs[total_header] = 0
|
||||||
|
|
||||||
if isinstance(payload, dict):
|
if isinstance(payload, dict):
|
||||||
# Read strings that count towards command length
|
# Read strings that count towards command length
|
||||||
# String length is length of longest localisation, including default.
|
# String length is length of longest localisation, including default.
|
||||||
for key in ('name', 'description', 'value'):
|
for key in ("name", "description", "value"):
|
||||||
if key in payload:
|
if key in payload:
|
||||||
value = payload[key]
|
value = payload[key]
|
||||||
if isinstance(value, str):
|
if isinstance(value, str):
|
||||||
values = (value, *payload.get(key + '_localizations', {}).values())
|
values = (value, *payload.get(key + "_localizations", {}).values())
|
||||||
maxlen = max(map(len, values))
|
maxlen = max(map(len, values))
|
||||||
total += maxlen
|
total += maxlen
|
||||||
breadcrumbs[(*header, key)] = maxlen
|
breadcrumbs[(*header, key)] = maxlen
|
||||||
@@ -824,7 +873,7 @@ def _recurse_length(payload, breadcrumbs={}, header=()) -> int:
|
|||||||
total += _recurse_length(value, breadcrumbs, loc)
|
total += _recurse_length(value, breadcrumbs, loc)
|
||||||
elif isinstance(payload, list):
|
elif isinstance(payload, list):
|
||||||
for i, item in enumerate(payload):
|
for i, item in enumerate(payload):
|
||||||
if isinstance(item, dict) and 'name' in item:
|
if isinstance(item, dict) and "name" in item:
|
||||||
loc = (*header, f"{i}<{item['name']}>")
|
loc = (*header, f"{i}<{item['name']}>")
|
||||||
else:
|
else:
|
||||||
loc = (*header, str(i))
|
loc = (*header, str(i))
|
||||||
@@ -837,11 +886,160 @@ def _recurse_length(payload, breadcrumbs={}, header=()) -> int:
|
|||||||
|
|
||||||
return total
|
return total
|
||||||
|
|
||||||
|
|
||||||
def write_records(records: list[dict[str, Any]], stream: StringIO):
|
def write_records(records: list[dict[str, Any]], stream: StringIO):
|
||||||
if records:
|
if records:
|
||||||
keys = records[0].keys()
|
keys = records[0].keys()
|
||||||
stream.write(','.join(keys))
|
stream.write(",".join(keys))
|
||||||
stream.write('\n')
|
stream.write("\n")
|
||||||
for record in records:
|
for record in records:
|
||||||
stream.write(','.join(map(str, record.values())))
|
stream.write(",".join(map(str, record.values())))
|
||||||
stream.write('\n')
|
stream.write("\n")
|
||||||
|
|
||||||
|
|
||||||
|
parse_dur_exps = [
|
||||||
|
(
|
||||||
|
r"(?P<value>\d+)\s*(?:(d)|(day))",
|
||||||
|
60 * 60 * 24,
|
||||||
|
),
|
||||||
|
(r"(?P<value>\d+)\s*(?:(h)|(hour))", 60 * 60),
|
||||||
|
(r"(?P<value>\d+)\s*(?:(m)|(min))", 60),
|
||||||
|
(r"(?P<value>\d+)\s*(?:(s)|(sec))", 1),
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def parse_duration(string: str) -> Optional[int]:
|
||||||
|
seconds = 0
|
||||||
|
found = False
|
||||||
|
for expr, multiplier in parse_dur_exps:
|
||||||
|
match = re.search(expr, string, flags=re.IGNORECASE)
|
||||||
|
if match:
|
||||||
|
found = True
|
||||||
|
seconds += int(match.group("value")) * multiplier
|
||||||
|
|
||||||
|
return seconds if found else None
|
||||||
|
|
||||||
|
|
||||||
|
async def pager(ctx, pages, locked=True, start_at=0, add_cancel=False, **kwargs):
|
||||||
|
"""
|
||||||
|
Shows the user each page from the provided list `pages` one at a time,
|
||||||
|
providing reactions to page back and forth between pages.
|
||||||
|
This is done asynchronously, and returns after displaying the first page.
|
||||||
|
|
||||||
|
Parameters
|
||||||
|
----------
|
||||||
|
pages: List(Union(str, discord.Embed))
|
||||||
|
A list of either strings or embeds to display as the pages.
|
||||||
|
locked: bool
|
||||||
|
Whether only the `ctx.author` should be able to use the paging reactions.
|
||||||
|
kwargs: ...
|
||||||
|
Remaining keyword arguments are transparently passed to the reply context method.
|
||||||
|
|
||||||
|
Returns: discord.Message
|
||||||
|
This is the output message, returned for easy deletion.
|
||||||
|
"""
|
||||||
|
cancel_emoji = cross
|
||||||
|
# Handle broken input
|
||||||
|
if len(pages) == 0:
|
||||||
|
raise ValueError("Pager cannot page with no pages!")
|
||||||
|
|
||||||
|
# Post first page. Method depends on whether the page is an embed or not.
|
||||||
|
if isinstance(pages[start_at], discord.Embed):
|
||||||
|
out_msg = await ctx.reply(embed=pages[start_at], **kwargs)
|
||||||
|
else:
|
||||||
|
out_msg = await ctx.reply(pages[start_at], **kwargs)
|
||||||
|
|
||||||
|
# Run the paging loop if required
|
||||||
|
if len(pages) > 1:
|
||||||
|
task = asyncio.create_task(
|
||||||
|
_pager(ctx, out_msg, pages, locked, start_at, add_cancel, **kwargs)
|
||||||
|
)
|
||||||
|
# ctx.tasks.append(task)
|
||||||
|
elif add_cancel:
|
||||||
|
await out_msg.add_reaction(cancel_emoji)
|
||||||
|
|
||||||
|
# Return the output message
|
||||||
|
return out_msg
|
||||||
|
|
||||||
|
|
||||||
|
async def _pager(ctx, out_msg, pages, locked, start_at, add_cancel, **kwargs):
|
||||||
|
"""
|
||||||
|
Asynchronous initialiser and loop for the `pager` utility above.
|
||||||
|
"""
|
||||||
|
# Page number
|
||||||
|
page = start_at
|
||||||
|
|
||||||
|
# Add reactions to the output message
|
||||||
|
next_emoji = "▶"
|
||||||
|
prev_emoji = "◀"
|
||||||
|
cancel_emoji = cross
|
||||||
|
|
||||||
|
try:
|
||||||
|
await out_msg.add_reaction(prev_emoji)
|
||||||
|
if add_cancel:
|
||||||
|
await out_msg.add_reaction(cancel_emoji)
|
||||||
|
await out_msg.add_reaction(next_emoji)
|
||||||
|
except discord.Forbidden:
|
||||||
|
# We don't have permission to add paging emojis
|
||||||
|
# Die as gracefully as we can
|
||||||
|
if ctx.guild:
|
||||||
|
perms = ctx.channel.permissions_for(ctx.guild.me)
|
||||||
|
if not perms.add_reactions:
|
||||||
|
await ctx.error_reply(
|
||||||
|
"Cannot page results because I do not have the `add_reactions` permission!"
|
||||||
|
)
|
||||||
|
elif not perms.read_message_history:
|
||||||
|
await ctx.error_reply(
|
||||||
|
"Cannot page results because I do not have the `read_message_history` permission!"
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
await ctx.error_reply(
|
||||||
|
"Cannot page results due to insufficient permissions!"
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
await ctx.error_reply("Cannot page results!")
|
||||||
|
return
|
||||||
|
|
||||||
|
# Check function to determine whether a reaction is valid
|
||||||
|
def check(reaction, user):
|
||||||
|
result = reaction.message.id == out_msg.id
|
||||||
|
result = result and str(reaction.emoji) in [next_emoji, prev_emoji]
|
||||||
|
result = result and not (user.id == ctx.bot.user.id)
|
||||||
|
result = result and not (locked and user != ctx.author)
|
||||||
|
return result
|
||||||
|
|
||||||
|
# Begin loop
|
||||||
|
while True:
|
||||||
|
# Wait for a valid reaction, break if we time out
|
||||||
|
try:
|
||||||
|
reaction, user = await ctx.bot.wait_for(
|
||||||
|
"reaction_add", check=check, timeout=300
|
||||||
|
)
|
||||||
|
except asyncio.TimeoutError:
|
||||||
|
break
|
||||||
|
|
||||||
|
# Attempt to remove the user's reaction, silently ignore errors
|
||||||
|
asyncio.ensure_future(out_msg.remove_reaction(reaction.emoji, user))
|
||||||
|
|
||||||
|
# Change the page number
|
||||||
|
page += 1 if reaction.emoji == next_emoji else -1
|
||||||
|
page %= len(pages)
|
||||||
|
|
||||||
|
# Edit the message with the new page
|
||||||
|
active_page = pages[page]
|
||||||
|
if isinstance(active_page, discord.Embed):
|
||||||
|
await out_msg.edit(embed=active_page, **kwargs)
|
||||||
|
else:
|
||||||
|
await out_msg.edit(content=active_page, **kwargs)
|
||||||
|
|
||||||
|
# Clean up by removing the reactions
|
||||||
|
try:
|
||||||
|
await out_msg.clear_reactions()
|
||||||
|
except discord.Forbidden:
|
||||||
|
try:
|
||||||
|
await out_msg.remove_reaction(next_emoji, ctx.client.user)
|
||||||
|
await out_msg.remove_reaction(prev_emoji, ctx.client.user)
|
||||||
|
except discord.NotFound:
|
||||||
|
pass
|
||||||
|
except discord.NotFound:
|
||||||
|
pass
|
||||||
|
|||||||
@@ -1,8 +1,12 @@
|
|||||||
import asyncio
|
import asyncio
|
||||||
import logging
|
import logging
|
||||||
|
from babel.translator import LocalBabel
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
util_babel = LocalBabel('utils')
|
||||||
|
|
||||||
from .hooked import *
|
from .hooked import *
|
||||||
from .leo import *
|
from .leo import *
|
||||||
from .micros import *
|
from .micros import *
|
||||||
|
from .msgeditor import *
|
||||||
|
|||||||
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user