Add vision model support: auto-switch on image detection, visionmodel command

This commit is contained in:
2026-08-17 00:40:04 -05:00
parent 71d71b4773
commit b81f3429d4
+58 -3
View File
@@ -45,6 +45,7 @@ class Assistant(commands.Cog):
enabled=True, enabled=True,
chat_channel=None, chat_channel=None,
model="gemma4", model="gemma4",
vision_model="gemma4",
temperature=0.7, temperature=0.7,
system_prompt="You are a helpful assistant.", system_prompt="You are a helpful assistant.",
question_mode=False, question_mode=False,
@@ -96,7 +97,29 @@ class Assistant(commands.Cog):
for mention in (f"<@{self.bot.user.id}>", f"<@!{self.bot.user.id}>"): for mention in (f"<@{self.bot.user.id}>", f"<@!{self.bot.user.id}>"):
content = content.replace(mention, "").strip() content = content.replace(mention, "").strip()
if not content: # Collect image URLs from this message and any replied-to message
image_urls = []
for attachment in message.attachments:
if attachment.content_type and attachment.content_type.startswith("image/"):
image_urls.append(attachment.url)
# Check replied-to message for images
if message.reference and message.reference.message_id:
try:
ref_msg = await message.channel.fetch_message(message.reference.message_id)
for attachment in ref_msg.attachments:
if attachment.content_type and attachment.content_type.startswith("image/"):
image_urls.append(attachment.url)
# Also check embeds for image URLs
for embed in ref_msg.embeds:
if embed.image and embed.image.url:
image_urls.append(embed.image.url)
if embed.thumbnail and embed.thumbnail.url:
image_urls.append(embed.thumbnail.url)
except (discord.NotFound, discord.HTTPException):
pass
if not content and not image_urls:
return return
if question_mode and not content.rstrip().endswith("?"): if question_mode and not content.rstrip().endswith("?"):
@@ -105,7 +128,7 @@ class Assistant(commands.Cog):
return return
async with message.channel.typing(): async with message.channel.typing():
reply = await self._chat(guild, message.channel.id, content) reply = await self._chat(guild, message.channel.id, content, image_urls=image_urls)
if reply is None: if reply is None:
await message.reply("No response from the assistant.") await message.reply("No response from the assistant.")
@@ -126,7 +149,7 @@ class Assistant(commands.Cog):
# API # API
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
async def _chat(self, guild: discord.Guild, channel_id: int, user_msg: str) -> str | None: async def _chat(self, guild: discord.Guild, channel_id: int, user_msg: str, image_urls: list[str] | None = None) -> str | None:
api_key = await self.config.api_key() api_key = await self.config.api_key()
if not api_key or not self.session: if not api_key or not self.session:
return None return None
@@ -137,10 +160,26 @@ class Assistant(commands.Cog):
max_tokens = await self.config.guild(guild).max_tokens() max_tokens = await self.config.guild(guild).max_tokens()
reasoning_effort = await self.config.guild(guild).reasoning_effort() reasoning_effort = await self.config.guild(guild).reasoning_effort()
# Auto-switch to vision model when images are present
if image_urls:
vision_model = await self.config.guild(guild).vision_model()
if vision_model:
model = vision_model
key = (guild.id, channel_id) key = (guild.id, channel_id)
if key not in self.history: if key not in self.history:
self.history[key] = [] self.history[key] = []
# Build user message content (text-only or multimodal)
if image_urls:
# OpenAI vision format: content is an array of parts
user_content = []
if user_msg:
user_content.append({"type": "text", "text": user_msg})
for url in image_urls:
user_content.append({"type": "image_url", "image_url": {"url": url}})
self.history[key].append({"role": "user", "content": user_content})
else:
self.history[key].append({"role": "user", "content": user_msg}) self.history[key].append({"role": "user", "content": user_msg})
messages = [{"role": "system", "content": system_prompt}] + self.history[key][-20:] messages = [{"role": "system", "content": system_prompt}] + self.history[key][-20:]
@@ -231,6 +270,20 @@ class Assistant(commands.Cog):
await self.config.guild(ctx.guild).model.set(model) await self.config.guild(ctx.guild).model.set(model)
await ctx.send(f"Model set to **{model}**.") await ctx.send(f"Model set to **{model}**.")
@assistant.command(name="visionmodel")
async def assistant_visionmodel(self, ctx: commands.Context, model: str = None):
"""Set the model used when images are detected. Omit to show current.
Vision-capable models: gemma4, green-l, green-r, gpt-oss-120b,
gemma-3-27b-it, mistral-small-3.2-24b-instruct-2506, kimi-k2.6, kimi-k3
"""
if model is None:
current = await self.config.guild(ctx.guild).vision_model()
await ctx.send(f"Vision model: **{current}**")
else:
await self.config.guild(ctx.guild).vision_model.set(model)
await ctx.send(f"Vision model set to **{model}**.")
@assistant.command(name="models") @assistant.command(name="models")
async def assistant_models(self, ctx: commands.Context): async def assistant_models(self, ctx: commands.Context):
"""List available chat models sorted by cost (cheapest first).""" """List available chat models sorted by cost (cheapest first)."""
@@ -414,6 +467,7 @@ class Assistant(commands.Cog):
enabled = await cfg.enabled() enabled = await cfg.enabled()
chat_channel = await cfg.chat_channel() chat_channel = await cfg.chat_channel()
model = await cfg.model() model = await cfg.model()
vision_model = await cfg.vision_model()
temp = await cfg.temperature() temp = await cfg.temperature()
sys_prompt = await cfg.system_prompt() sys_prompt = await cfg.system_prompt()
q_mode = await cfg.question_mode() q_mode = await cfg.question_mode()
@@ -428,6 +482,7 @@ class Assistant(commands.Cog):
embed.add_field(name="Chat Channel", value=channel_str, inline=True) embed.add_field(name="Chat Channel", value=channel_str, inline=True)
embed.add_field(name="API Key", value="Set" if has_key else "Not set", inline=True) embed.add_field(name="API Key", value="Set" if has_key else "Not set", inline=True)
embed.add_field(name="Model", value=model, inline=True) embed.add_field(name="Model", value=model, inline=True)
embed.add_field(name="Vision Model", value=vision_model, inline=True)
embed.add_field(name="Temperature", value=str(temp), inline=True) embed.add_field(name="Temperature", value=str(temp), inline=True)
embed.add_field(name="Question Mode", value=str(q_mode), inline=True) embed.add_field(name="Question Mode", value=str(q_mode), inline=True)
embed.add_field(name="Max Tokens", value=str(max_tokens) if max_tokens else "Unlimited", inline=True) embed.add_field(name="Max Tokens", value=str(max_tokens) if max_tokens else "Unlimited", inline=True)