forked from jakobdylanc/llmcord
-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathllmcord.py
More file actions
360 lines (262 loc) · 15.5 KB
/
Copy pathllmcord.py
File metadata and controls
360 lines (262 loc) · 15.5 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
360
import asyncio
from base64 import b64encode
from dataclasses import dataclass, field
from datetime import datetime
import logging
from typing import Any, Literal, Optional
import discord
from discord.app_commands import Choice
from discord.ext import commands
from discord.ui import LayoutView, TextDisplay
import httpx
from openai import AsyncOpenAI
import yaml
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s %(levelname)s: %(message)s",
)
VISION_MODEL_TAGS = ("claude", "gemini", "gemma", "gpt-4", "gpt-5", "grok-4", "llama", "llava", "mistral", "o3", "o4", "vision", "vl")
PROVIDERS_SUPPORTING_USERNAMES = ("openai", "x-ai")
EMBED_COLOR_COMPLETE = discord.Color.dark_green()
EMBED_COLOR_INCOMPLETE = discord.Color.orange()
STREAMING_INDICATOR = " ⚪"
EDIT_DELAY_SECONDS = 1
MAX_MESSAGE_NODES = 500
def get_config(filename: str = "config.yaml") -> dict[str, Any]:
with open(filename, encoding="utf-8") as file:
return yaml.safe_load(file)
config = get_config()
curr_model = next(iter(config["models"]))
msg_nodes = {}
last_task_time = 0
intents = discord.Intents.default()
intents.message_content = True
activity = discord.CustomActivity(name=(config.get("status_message") or "github.com/jakobdylanc/llmcord")[:128])
discord_bot = commands.Bot(intents=intents, activity=activity, command_prefix=None)
httpx_client = httpx.AsyncClient()
@dataclass
class MsgNode:
text: Optional[str] = None
images: list[dict[str, Any]] = field(default_factory=list)
role: Literal["user", "assistant"] = "assistant"
user_id: Optional[int] = None
has_bad_attachments: bool = False
fetch_parent_failed: bool = False
parent_msg: Optional[discord.Message] = None
lock: asyncio.Lock = field(default_factory=asyncio.Lock)
@discord_bot.tree.command(name="model", description="View or switch the current model")
async def model_command(interaction: discord.Interaction, model: str) -> None:
global curr_model
if model == curr_model:
output = f"Current model: `{curr_model}`"
else:
if user_is_admin := interaction.user.id in config["permissions"]["users"]["admin_ids"]:
curr_model = model
output = f"Model switched to: `{model}`"
logging.info(output)
else:
output = "You don't have permission to change the model."
await interaction.response.send_message(output, ephemeral=(interaction.channel.type == discord.ChannelType.private))
@model_command.autocomplete("model")
async def model_autocomplete(interaction: discord.Interaction, curr_str: str) -> list[Choice[str]]:
global config
if curr_str == "":
config = await asyncio.to_thread(get_config)
choices = [Choice(name=f"◉ {curr_model} (current)", value=curr_model)] if curr_str.lower() in curr_model.lower() else []
choices += [Choice(name=f"○ {model}", value=model) for model in config["models"] if model != curr_model and curr_str.lower() in model.lower()]
return choices[:25]
@discord_bot.event
async def on_ready() -> None:
if client_id := config.get("client_id"):
logging.info(f"\n\nBOT INVITE URL:\nhttps://discord.com/oauth2/authorize?client_id={client_id}&permissions=412317191168&scope=bot\n")
await discord_bot.tree.sync()
@discord_bot.event
async def on_message(new_msg: discord.Message) -> None:
global last_task_time
is_dm = new_msg.channel.type == discord.ChannelType.private
if new_msg.author.bot:
return
role_ids = set(role.id for role in getattr(new_msg.author, "roles", ()))
channel_ids = set(filter(None, (new_msg.channel.id, getattr(new_msg.channel, "parent_id", None), getattr(new_msg.channel, "category_id", None))))
config = await asyncio.to_thread(get_config)
allow_dms = config.get("allow_dms", True)
permissions = config["permissions"]
user_is_admin = new_msg.author.id in permissions["users"]["admin_ids"]
# Get permission lists
allowed_user_ids = permissions["users"].get("allowed_ids", [])
blocked_user_ids = permissions["users"].get("blocked_ids", [])
allowed_role_ids = permissions["roles"].get("allowed_ids", [])
blocked_role_ids = permissions["roles"].get("blocked_ids", [])
allowed_channel_ids = permissions["channels"].get("allowed_ids", [])
blocked_channel_ids = permissions["channels"].get("blocked_ids", [])
# Check explicit allowlists (for mention bypass only)
user_in_allowlist = user_is_admin or new_msg.author.id in allowed_user_ids or any(id in allowed_role_ids for id in role_ids)
channel_in_allowlist = user_is_admin or any(id in allowed_channel_ids for id in channel_ids)
# MENTION LOGIC:
# - Skip mention requirement if: DM, or channel in allowlist, or user in allowlist
# - Otherwise require mention
if not is_dm and not channel_in_allowlist and not user_in_allowlist:
if discord_bot.user not in new_msg.mentions:
return
# PERMISSION LOGIC: Everyone is allowed unless blocked
# Check if user is blocked
user_blocked = new_msg.author.id in blocked_user_ids or any(id in blocked_role_ids for id in role_ids)
# Check if channel is blocked (DMs check allow_dms config)
if is_dm:
channel_blocked = not allow_dms
else:
channel_blocked = any(id in blocked_channel_ids for id in channel_ids)
if user_blocked or channel_blocked:
return
# ... rest of the function remains exactly the same ...
provider_slash_model = curr_model
provider, model = provider_slash_model.removesuffix(":vision").split("/", 1)
provider_config = config["providers"][provider]
base_url = provider_config["base_url"]
api_key = provider_config.get("api_key", "sk-no-key-required")
openai_client = AsyncOpenAI(base_url=base_url, api_key=api_key)
model_parameters = config["models"].get(provider_slash_model, None)
extra_headers = provider_config.get("extra_headers")
extra_query = provider_config.get("extra_query")
extra_body = (provider_config.get("extra_body") or {}) | (model_parameters or {}) or None
accept_images = any(x in provider_slash_model.lower() for x in VISION_MODEL_TAGS)
accept_usernames = any(provider_slash_model.lower().startswith(x) for x in PROVIDERS_SUPPORTING_USERNAMES)
max_text = config.get("max_text", 100000)
max_images = config.get("max_images", 5) if accept_images else 0
max_messages = config.get("max_messages", 25)
# Build message chain and set user warnings
messages = []
user_warnings = set()
curr_msg = new_msg
while curr_msg != None and len(messages) < max_messages:
curr_node = msg_nodes.setdefault(curr_msg.id, MsgNode())
async with curr_node.lock:
if curr_node.text == None:
cleaned_content = curr_msg.content.removeprefix(discord_bot.user.mention).lstrip()
good_attachments = [att for att in curr_msg.attachments if att.content_type and any(att.content_type.startswith(x) for x in ("text", "image"))]
attachment_responses = await asyncio.gather(*[httpx_client.get(att.url) for att in good_attachments])
curr_node.text = "\n".join(
([cleaned_content] if cleaned_content else [])
+ ["\n".join(filter(None, (embed.title, embed.description, embed.footer.text))) for embed in curr_msg.embeds]
+ [component.content for component in curr_msg.components if component.type == discord.ComponentType.text_display]
+ [resp.text for att, resp in zip(good_attachments, attachment_responses) if att.content_type.startswith("text")]
)
curr_node.images = [
dict(type="image_url", image_url=dict(url=f"data:{att.content_type};base64,{b64encode(resp.content).decode('utf-8')}"))
for att, resp in zip(good_attachments, attachment_responses)
if att.content_type.startswith("image")
]
curr_node.role = "assistant" if curr_msg.author == discord_bot.user else "user"
curr_node.user_id = curr_msg.author.id if curr_node.role == "user" else None
curr_node.has_bad_attachments = len(curr_msg.attachments) > len(good_attachments)
try:
if (
curr_msg.reference == None
and discord_bot.user.mention not in curr_msg.content
and (prev_msg_in_channel := ([m async for m in curr_msg.channel.history(before=curr_msg, limit=1)] or [None])[0])
and prev_msg_in_channel.type in (discord.MessageType.default, discord.MessageType.reply)
and prev_msg_in_channel.author == (discord_bot.user if curr_msg.channel.type == discord.ChannelType.private else curr_msg.author)
):
curr_node.parent_msg = prev_msg_in_channel
else:
is_public_thread = curr_msg.channel.type == discord.ChannelType.public_thread
parent_is_thread_start = is_public_thread and curr_msg.reference == None and curr_msg.channel.parent.type == discord.ChannelType.text
if parent_msg_id := curr_msg.channel.id if parent_is_thread_start else getattr(curr_msg.reference, "message_id", None):
if parent_is_thread_start:
curr_node.parent_msg = curr_msg.channel.starter_message or await curr_msg.channel.parent.fetch_message(parent_msg_id)
else:
curr_node.parent_msg = curr_msg.reference.cached_message or await curr_msg.channel.fetch_message(parent_msg_id)
except (discord.NotFound, discord.HTTPException):
logging.exception("Error fetching next message in the chain")
curr_node.fetch_parent_failed = True
if curr_node.images[:max_images]:
content = ([dict(type="text", text=curr_node.text[:max_text])] if curr_node.text[:max_text] else []) + curr_node.images[:max_images]
else:
content = curr_node.text[:max_text]
if content != "":
message = dict(content=content, role=curr_node.role)
if accept_usernames and curr_node.user_id != None:
message["name"] = str(curr_node.user_id)
messages.append(message)
if len(curr_node.text) > max_text:
user_warnings.add(f"⚠️ Max {max_text:,} characters per message")
if len(curr_node.images) > max_images:
user_warnings.add(f"⚠️ Max {max_images} image{'' if max_images == 1 else 's'} per message" if max_images > 0 else "⚠️ Can't see images")
if curr_node.has_bad_attachments:
user_warnings.add("⚠️ Unsupported attachments")
if curr_node.fetch_parent_failed or (curr_node.parent_msg != None and len(messages) == max_messages):
user_warnings.add(f"⚠️ Only using last {len(messages)} message{'' if len(messages) == 1 else 's'}")
curr_msg = curr_node.parent_msg
logging.info(f"Message received (user ID: {new_msg.author.id}, attachments: {len(new_msg.attachments)}, conversation length: {len(messages)}):\n{new_msg.content}")
if system_prompt := config.get("system_prompt"):
now = datetime.now().astimezone()
system_prompt = system_prompt.replace("{date}", now.strftime("%B %d %Y")).replace("{time}", now.strftime("%H:%M:%S %Z%z")).strip()
if accept_usernames:
system_prompt += "\n\nUser's names are their Discord IDs and should be typed as '<@ID>'."
messages.append(dict(role="system", content=system_prompt))
# Generate and send response message(s) (can be multiple if response is long)
curr_content = finish_reason = None
response_msgs = []
response_contents = []
openai_kwargs = dict(model=model, messages=messages[::-1], stream=True, extra_headers=extra_headers, extra_query=extra_query, extra_body=extra_body)
if use_plain_responses := config.get("use_plain_responses", False):
max_message_length = 4000
else:
max_message_length = 4096 - len(STREAMING_INDICATOR)
embed = discord.Embed.from_dict(dict(fields=[dict(name=warning, value="", inline=False) for warning in sorted(user_warnings)]))
async def reply_helper(**reply_kwargs) -> None:
reply_target = new_msg if not response_msgs else response_msgs[-1]
response_msg = await reply_target.reply(**reply_kwargs)
response_msgs.append(response_msg)
msg_nodes[response_msg.id] = MsgNode(parent_msg=new_msg)
await msg_nodes[response_msg.id].lock.acquire()
try:
async with new_msg.channel.typing():
async for chunk in await openai_client.chat.completions.create(**openai_kwargs):
if finish_reason != None:
break
if not (choice := chunk.choices[0] if chunk.choices else None):
continue
finish_reason = choice.finish_reason
prev_content = curr_content or ""
curr_content = choice.delta.content or ""
new_content = prev_content if finish_reason == None else (prev_content + curr_content)
if response_contents == [] and new_content == "":
continue
if start_next_msg := response_contents == [] or len(response_contents[-1] + new_content) > max_message_length:
response_contents.append("")
response_contents[-1] += new_content
if not use_plain_responses:
time_delta = datetime.now().timestamp() - last_task_time
ready_to_edit = time_delta >= EDIT_DELAY_SECONDS
msg_split_incoming = finish_reason == None and len(response_contents[-1] + curr_content) > max_message_length
is_final_edit = finish_reason != None or msg_split_incoming
is_good_finish = finish_reason != None and finish_reason.lower() in ("stop", "end_turn")
if start_next_msg or ready_to_edit or is_final_edit:
embed.description = response_contents[-1] if is_final_edit else (response_contents[-1] + STREAMING_INDICATOR)
embed.color = EMBED_COLOR_COMPLETE if msg_split_incoming or is_good_finish else EMBED_COLOR_INCOMPLETE
if start_next_msg:
await reply_helper(embed=embed, silent=True)
else:
await asyncio.sleep(EDIT_DELAY_SECONDS - time_delta)
await response_msgs[-1].edit(embed=embed)
last_task_time = datetime.now().timestamp()
if use_plain_responses:
for content in response_contents:
await reply_helper(view=LayoutView().add_item(TextDisplay(content=content)))
except Exception:
logging.exception("Error while generating response")
for response_msg in response_msgs:
msg_nodes[response_msg.id].text = "".join(response_contents)
msg_nodes[response_msg.id].lock.release()
# Delete oldest MsgNodes (lowest message IDs) from the cache
if (num_nodes := len(msg_nodes)) > MAX_MESSAGE_NODES:
for msg_id in sorted(msg_nodes.keys())[: num_nodes - MAX_MESSAGE_NODES]:
async with msg_nodes.setdefault(msg_id, MsgNode()).lock:
msg_nodes.pop(msg_id, None)
async def main() -> None:
await discord_bot.start(config["bot_token"])
try:
asyncio.run(main())
except KeyboardInterrupt:
pass