-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathmain.py
More file actions
359 lines (283 loc) · 13.6 KB
/
Copy pathmain.py
File metadata and controls
359 lines (283 loc) · 13.6 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
import asyncio
import traceback
from pathlib import Path
import aiohttp
import click
import discord
import logging
import os
import typing as tp
import uvicorn
from discord import InteractionType
from discord.ext import commands
from discord.utils import _ColourFormatter
from fastapi import FastAPI, Depends, HTTPException
from fastapi.security import APIKeyHeader
from pydantic import BaseModel, Field as PydanticField
from components.modals.approval import ApprovalModal
from components.modals.pre_approval import PreApprovalModal
from components.modals.pre_rejection import PreRejectionModal
from components.modals.pre_rejection_no_review import PreRejectionNoReviewModal
from components.modals.rejection import RejectionModal
from components.modals.request_submission import RequestSubmissionModal
from components.modals.trainee_review_feedback import TraineeReviewFeedbackModal
from components.views.pending_request_widget import PendingRequestWidgetApproveAndReviewBtn, PendingRequestWidgetJustApproveBtn, PendingRequestWidgetJustRejectBtn, PendingRequestWidgetRejectAndReviewBtn
from components.views.resolution_widget import ResolutionWidgetEpicBtn, ResolutionWidgetFeatureBtn, ResolutionWidgetLegendaryBtn, ResolutionWidgetMythicBtn, ResolutionWidgetRejectBtn, ResolutionWidgetStarrateBtn
from components.views.trainee_pick_widget import TraineePickWidgetAcceptBtn, TraineePickWidgetRejectBtn
from components.views.trainee_promotion_decision import TraineePromotionDecisionExpelBtn, TraineePromotionDecisionPromoteBtn, TraineePromotionDecisionWaitBtn
from components.views.trainee_review_widget import TraineeReviewWidgetAcceptBtn, TraineeReviewWidgetRejectBtn
from config.texts import validate as validate_texts
from config.routes import validate as validate_routes
from config.parameters import validate as validate_parameters
from config.stage_parameters import validate as validate_stage_parameters, get_value as get_stage_parameter_value
from config.permission_flags import validate as validate_permission_flags
from db import EngineProvider
from db.models import RouteID, Request
from util.datatypes import Language, Opinion
from facades.reports import stream_results_chart, StreamResolution
from facades.requests import add_opinion, complete_request, create_limbo_request, get_existing_opinion, get_latest_pending_request, get_pending_request, is_request_unresolved, resolve
from globalconf import CONFIG
from services.disc import post_raw_text
from util.datatypes import SendType, Stage
from util.identifiers import StageParameterID
from util.translator import Translator
class StreamAnnouncementPayload(BaseModel):
text: str = PydanticField(..., min_length=1, max_length=2000)
class StreamGoodbyePayload(BaseModel):
text: str = PydanticField(..., min_length=1, max_length=2000)
not_reviewed: int
approved: int
rejected: int
later: int
class RequestResolutionPayload(BaseModel):
request_id: int
sent_for: SendType | None = None
stream_link: str | None = None
class RequestCreationPayload(BaseModel):
level_id: int
creator_name: str
language: Language
showcase_yt_link: str
class RequestPreApprovalPayload(BaseModel):
request_id: int
api_app = FastAPI()
header_scheme = APIKeyHeader(name="x-key")
@api_app.post("/message/stream_start")
async def send_stream_start_message(payload: StreamAnnouncementPayload, key: str = Depends(header_scheme)) -> None:
if key != os.getenv("API_TOKEN"):
raise HTTPException(status_code=401, detail="Wrong token")
await post_raw_text(RouteID.STREAM_START_ANNOUNCEMENT, payload.text)
@api_app.post("/message/stream_end")
async def send_stream_end_message(payload: StreamGoodbyePayload, key: str = Depends(header_scheme)) -> None:
if key != os.getenv("API_TOKEN"):
raise HTTPException(status_code=401, detail="Wrong token")
report_path = None
if payload.approved + payload.rejected + payload.later: # If at least one request was reviewed
report_path = stream_results_chart(counts={
StreamResolution.APPROVED: payload.approved,
StreamResolution.REJECTED: payload.rejected,
StreamResolution.LATER: payload.later,
StreamResolution.NOT_REVIEWED: payload.not_reviewed
})
await post_raw_text(RouteID.STREAM_END_GOODBYE, payload.text, file_path=report_path)
Path(report_path).unlink(missing_ok=True)
@api_app.get("/request/random")
async def random_request(key: str = Depends(header_scheme)) -> Request:
if key != os.getenv("API_TOKEN"):
raise HTTPException(status_code=401, detail="Wrong token")
return await get_pending_request(oldest=False)
@api_app.get("/request/oldest")
async def oldest_request(key: str = Depends(header_scheme)) -> Request:
if key != os.getenv("API_TOKEN"):
raise HTTPException(status_code=401, detail="Wrong token")
return await get_pending_request(oldest=True)
@api_app.post("/request/resolve")
async def request_resolve(payload: RequestResolutionPayload, key: str = Depends(header_scheme)) -> bool:
if key != os.getenv("API_TOKEN"):
raise HTTPException(status_code=401, detail="Wrong token")
if not await is_request_unresolved(payload.request_id):
return False
if payload.stream_link:
reason = f"Reviewed on stream: {payload.stream_link}"
else:
reason = "Reviewed on stream"
return await resolve(
resolving_mod=CONFIG.admin,
request_id=payload.request_id,
sent_for=payload.sent_for,
review_text=None,
reason=reason
)
@api_app.post("/request/preapprove")
async def request_pre_approve(payload: RequestPreApprovalPayload, key: str = Depends(header_scheme)) -> None:
if key != os.getenv("API_TOKEN"):
raise HTTPException(status_code=401, detail="Wrong token")
existing_opinion = await get_existing_opinion(reviewer=CONFIG.admin, request_id=payload.request_id, resolution_only=False)
if not existing_opinion:
await add_opinion(
reviewer=CONFIG.admin,
request_id=payload.request_id,
opinion=Opinion.APPROVED,
review_text=None,
reason="Marked as 'Later' on stream"
)
async def create_single_request(payload: RequestCreationPayload, allow_queue_closing: bool) -> int:
existing_request = await get_latest_pending_request(payload.level_id)
if existing_request:
return existing_request.id
request_id = await create_limbo_request(
level_id=payload.level_id,
request_language=payload.language,
invoker=CONFIG.admin,
creator=payload.creator_name
)
await complete_request(
request_id=request_id,
yt_link=payload.showcase_yt_link,
additional_comment="Requested on stream",
invoker=CONFIG.admin,
allow_queue_closing=allow_queue_closing
)
return request_id
@api_app.post("/request/create")
async def request_create(payload: RequestCreationPayload, key: str = Depends(header_scheme)) -> int:
if key != os.getenv("API_TOKEN"):
raise HTTPException(status_code=401, detail="Wrong token")
return await create_single_request(payload, False)
@api_app.post("/request/create_batch")
async def request_create_batch(payload: list[RequestCreationPayload], key: str = Depends(header_scheme)) -> None:
if key != os.getenv("API_TOKEN"):
raise HTTPException(status_code=401, detail="Wrong token")
for single_request_payload in payload:
await create_single_request(single_request_payload, True)
class RequestBot(commands.Bot):
client: aiohttp.ClientSession
def __init__(self, *args: tp.Any, **kwargs: tp.Any) -> None:
self.logger = logging.getLogger(self.__class__.__name__)
handler = logging.StreamHandler()
formatter = _ColourFormatter()
handler.setFormatter(formatter)
self.logger.addHandler(handler)
self.logger.setLevel(logging.DEBUG)
self.ext_dir = "cogs"
self.synced = False
self.guild_id = 0
intents = discord.Intents.default()
intents.message_content = True
super().__init__(*args, **kwargs, command_prefix=commands.when_mentioned, intents=intents) # noqa
async def on_ready(self) -> None:
self.logger.info(f"Logged in as {self.user} ({self.user.id})")
await self.change_presence(activity=discord.Activity(type=discord.ActivityType.watching, name="out for your level requests!"))
CONFIG.guild = self.get_guild(self.guild_id)
assert CONFIG.guild
CONFIG.admin = await CONFIG.guild.fetch_member(get_stage_parameter_value(StageParameterID.ADMIN_USER_ID))
await EngineProvider.load()
async def _load_extensions(self) -> None:
if not os.path.isdir(self.ext_dir):
self.logger.error(f"Extension directory {self.ext_dir} does not exist.")
return
for filename in os.listdir(self.ext_dir):
if filename.endswith(".py") and not filename.startswith("_"):
try:
await self.load_extension(f"{self.ext_dir}.{filename[:-3]}")
self.logger.info(f"Loaded extension {filename[:-3]}")
except commands.ExtensionError:
self.logger.error(f"Failed to load extension {filename[:-3]}\n{traceback.format_exc()}")
async def on_error(self, event_method: str, *args: tp.Any, **kwargs: tp.Any) -> None:
self.logger.error(f"An error occurred in {event_method}.\n{traceback.format_exc()}")
async def close(self) -> None:
await super().close()
await self.client.close()
async def sync_tree(self) -> None:
guild = discord.Object(id=self.guild_id)
if not self.synced:
self.tree.copy_global_to(guild=guild)
self.logger.info("Tree copied")
await self.tree.set_translator(Translator())
self.logger.info("Translator set")
result = await self.tree.sync(guild=guild)
self.logger.info(f"Synced command tree: {len(result)} cogs")
self.synced = not self.synced
else:
self.logger.info("Skipped syncing command tree")
async def setup_hook(self) -> None:
self.client = aiohttp.ClientSession()
await self._load_extensions()
self.guild_id = get_stage_parameter_value(StageParameterID.GUILD_ID)
self.add_dynamic_items(PendingRequestWidgetApproveAndReviewBtn)
self.add_dynamic_items(PendingRequestWidgetRejectAndReviewBtn)
self.add_dynamic_items(PendingRequestWidgetJustApproveBtn)
self.add_dynamic_items(PendingRequestWidgetJustRejectBtn)
self.add_dynamic_items(ResolutionWidgetStarrateBtn)
self.add_dynamic_items(ResolutionWidgetFeatureBtn)
self.add_dynamic_items(ResolutionWidgetEpicBtn)
self.add_dynamic_items(ResolutionWidgetMythicBtn)
self.add_dynamic_items(ResolutionWidgetLegendaryBtn)
self.add_dynamic_items(ResolutionWidgetRejectBtn)
self.add_dynamic_items(TraineeReviewWidgetAcceptBtn)
self.add_dynamic_items(TraineeReviewWidgetRejectBtn)
self.add_dynamic_items(TraineePromotionDecisionPromoteBtn)
self.add_dynamic_items(TraineePromotionDecisionExpelBtn)
self.add_dynamic_items(TraineePromotionDecisionWaitBtn)
self.add_dynamic_items(TraineePickWidgetAcceptBtn)
self.add_dynamic_items(TraineePickWidgetRejectBtn)
self.logger.info("Added dynamic items")
@staticmethod
async def on_interaction(inter: discord.Interaction):
custom_id = inter.data.get("custom_id")
if custom_id and inter.type == InteractionType.modal_submit:
if custom_id.startswith("rsm:"):
await RequestSubmissionModal.handle_interaction(inter)
elif custom_id.startswith("prnrm:"):
await PreRejectionNoReviewModal.handle_interaction(inter)
elif custom_id.startswith("prm:"):
await PreRejectionModal.handle_interaction(inter)
elif custom_id.startswith("pam:"):
await PreApprovalModal.handle_interaction(inter)
elif custom_id.startswith("am:"):
await ApprovalModal.handle_interaction(inter)
elif custom_id.startswith("rm:"):
await RejectionModal.handle_interaction(inter)
elif custom_id.startswith("trf:"):
await TraineeReviewFeedbackModal.handle_interaction(inter)
def start(self, *args: tp.Any, **kwargs: tp.Any) -> tp.Coroutine[tp.Any, tp.Any, None]:
try:
return super().start(os.getenv("BOT_TOKEN"), *args, **kwargs)
except (discord.LoginFailure, KeyboardInterrupt):
self.logger.info("Exiting...")
exit()
async def start_api() -> None:
config = uvicorn.Config(api_app, host="0.0.0.0", port=5000, log_level="info")
server = uvicorn.Server(config)
await server.serve()
@click.command
@click.option(
"--debug",
is_flag=True,
help="Use testing server"
)
@click.option(
"--log_queries",
is_flag=True,
help="Log queries to the output"
)
def main(debug: bool, log_queries: bool) -> None:
if log_queries:
logger = logging.getLogger('sqlalchemy.engine')
logger.setLevel(logging.DEBUG)
CONFIG.stage = Stage.TEST if debug else Stage.PROD
validate_texts()
validate_routes()
validate_parameters()
validate_stage_parameters()
validate_permission_flags()
loop = asyncio.new_event_loop()
asyncio.set_event_loop(loop)
CONFIG.bot = RequestBot()
tasks = [
loop.create_task(CONFIG.bot.start()),
loop.create_task(start_api()),
]
loop.run_until_complete(asyncio.wait(tasks))
if __name__ == "__main__":
main()