mirror of
https://github.com/jxxghp/MoviePilot.git
synced 2026-09-05 23:47:41 +08:00
feat(agent): 支持图片消息处理
This commit is contained in:
+55
-40
@@ -10,7 +10,7 @@ from langchain.agents.middleware import (
|
|||||||
SummarizationMiddleware,
|
SummarizationMiddleware,
|
||||||
LLMToolSelectorMiddleware,
|
LLMToolSelectorMiddleware,
|
||||||
)
|
)
|
||||||
from langchain_core.messages import (
|
from langchain_core.messages import ( # noqa: F401
|
||||||
HumanMessage,
|
HumanMessage,
|
||||||
BaseMessage,
|
BaseMessage,
|
||||||
)
|
)
|
||||||
@@ -42,12 +42,12 @@ class MoviePilotAgent:
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
session_id: str,
|
session_id: str,
|
||||||
user_id: str = None,
|
user_id: str = None,
|
||||||
channel: str = None,
|
channel: str = None,
|
||||||
source: str = None,
|
source: str = None,
|
||||||
username: str = None,
|
username: str = None,
|
||||||
):
|
):
|
||||||
self.session_id = session_id
|
self.session_id = session_id
|
||||||
self.user_id = user_id
|
self.user_id = user_id
|
||||||
@@ -92,10 +92,10 @@ class MoviePilotAgent:
|
|||||||
if block.get("thought"):
|
if block.get("thought"):
|
||||||
continue
|
continue
|
||||||
if block.get("type") in (
|
if block.get("type") in (
|
||||||
"thinking",
|
"thinking",
|
||||||
"reasoning_content",
|
"reasoning_content",
|
||||||
"reasoning",
|
"reasoning",
|
||||||
"thought",
|
"thought",
|
||||||
):
|
):
|
||||||
continue
|
continue
|
||||||
if block.get("type") == "text":
|
if block.get("type") == "text":
|
||||||
@@ -174,20 +174,32 @@ class MoviePilotAgent:
|
|||||||
logger.error(f"创建 Agent 失败: {e}")
|
logger.error(f"创建 Agent 失败: {e}")
|
||||||
raise e
|
raise e
|
||||||
|
|
||||||
async def process(self, message: str) -> str:
|
async def process(self, message: str, images: List[str] = None) -> str:
|
||||||
"""
|
"""
|
||||||
处理用户消息,流式推理并返回 Agent 回复
|
处理用户消息,流式推理并返回 Agent 回复
|
||||||
"""
|
"""
|
||||||
try:
|
try:
|
||||||
logger.info(f"Agent推理: session_id={self.session_id}, input={message}")
|
logger.info(
|
||||||
|
f"Agent推理: session_id={self.session_id}, input={message}, images={len(images) if images else 0}"
|
||||||
|
)
|
||||||
|
|
||||||
# 获取历史消息
|
# 获取历史消息
|
||||||
messages = memory_manager.get_agent_messages(
|
messages = memory_manager.get_agent_messages(
|
||||||
session_id=self.session_id, user_id=self.user_id
|
session_id=self.session_id, user_id=self.user_id
|
||||||
)
|
)
|
||||||
|
|
||||||
# 增加用户消息
|
# 构建用户消息内容
|
||||||
messages.append(HumanMessage(content=message))
|
if images:
|
||||||
|
from langchain_core.messages import HumanMessage
|
||||||
|
|
||||||
|
content = []
|
||||||
|
if message:
|
||||||
|
content.append({"type": "text", "text": message})
|
||||||
|
for img in images:
|
||||||
|
content.append({"type": "image_url", "image_url": {"url": img}})
|
||||||
|
messages.append(HumanMessage(content=content))
|
||||||
|
else:
|
||||||
|
messages.append(HumanMessage(content=message))
|
||||||
|
|
||||||
# 执行推理
|
# 执行推理
|
||||||
await self._execute_agent(messages)
|
await self._execute_agent(messages)
|
||||||
@@ -199,7 +211,7 @@ class MoviePilotAgent:
|
|||||||
return error_message
|
return error_message
|
||||||
|
|
||||||
async def _stream_agent_tokens(
|
async def _stream_agent_tokens(
|
||||||
self, agent, messages: dict, config: dict, on_token: Callable[[str], None]
|
self, agent, messages: dict, config: dict, on_token: Callable[[str], None]
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
流式运行智能体,过滤工具调用token和思考内容,将模型生成的内容通过回调输出。
|
流式运行智能体,过滤工具调用token和思考内容,将模型生成的内容通过回调输出。
|
||||||
@@ -212,18 +224,18 @@ class MoviePilotAgent:
|
|||||||
buffer = ""
|
buffer = ""
|
||||||
|
|
||||||
async for chunk in agent.astream(
|
async for chunk in agent.astream(
|
||||||
messages,
|
messages,
|
||||||
stream_mode="messages",
|
stream_mode="messages",
|
||||||
config=config,
|
config=config,
|
||||||
subgraphs=False,
|
subgraphs=False,
|
||||||
version="v2",
|
version="v2",
|
||||||
):
|
):
|
||||||
if chunk["type"] == "messages":
|
if chunk["type"] == "messages":
|
||||||
token, metadata = chunk["data"]
|
token, metadata = chunk["data"]
|
||||||
if (
|
if (
|
||||||
token
|
token
|
||||||
and hasattr(token, "tool_call_chunks")
|
and hasattr(token, "tool_call_chunks")
|
||||||
and not token.tool_call_chunks
|
and not token.tool_call_chunks
|
||||||
):
|
):
|
||||||
# 跳过模型思考/推理内容(如 DeepSeek R1 的 reasoning_content)
|
# 跳过模型思考/推理内容(如 DeepSeek R1 的 reasoning_content)
|
||||||
additional = getattr(token, "additional_kwargs", None)
|
additional = getattr(token, "additional_kwargs", None)
|
||||||
@@ -241,7 +253,7 @@ class MoviePilotAgent:
|
|||||||
if start_idx > 0:
|
if start_idx > 0:
|
||||||
on_token(buffer[:start_idx])
|
on_token(buffer[:start_idx])
|
||||||
in_think_tag = True
|
in_think_tag = True
|
||||||
buffer = buffer[start_idx + 7:]
|
buffer = buffer[start_idx + 7 :]
|
||||||
else:
|
else:
|
||||||
# 检查是否以 <think> 的前缀结尾
|
# 检查是否以 <think> 的前缀结尾
|
||||||
partial_match = False
|
partial_match = False
|
||||||
@@ -259,7 +271,7 @@ class MoviePilotAgent:
|
|||||||
end_idx = buffer.find("</think>")
|
end_idx = buffer.find("</think>")
|
||||||
if end_idx != -1:
|
if end_idx != -1:
|
||||||
in_think_tag = False
|
in_think_tag = False
|
||||||
buffer = buffer[end_idx + 8:]
|
buffer = buffer[end_idx + 8 :]
|
||||||
else:
|
else:
|
||||||
# 检查是否以 </think> 的前缀结尾
|
# 检查是否以 </think> 的前缀结尾
|
||||||
partial_match = False
|
partial_match = False
|
||||||
@@ -421,6 +433,7 @@ class _MessageTask:
|
|||||||
session_id: str
|
session_id: str
|
||||||
user_id: str
|
user_id: str
|
||||||
message: str
|
message: str
|
||||||
|
images: Optional[List[str]] = None
|
||||||
channel: Optional[str] = None
|
channel: Optional[str] = None
|
||||||
source: Optional[str] = None
|
source: Optional[str] = None
|
||||||
username: Optional[str] = None
|
username: Optional[str] = None
|
||||||
@@ -467,13 +480,14 @@ class AgentManager:
|
|||||||
self.active_agents.clear()
|
self.active_agents.clear()
|
||||||
|
|
||||||
async def process_message(
|
async def process_message(
|
||||||
self,
|
self,
|
||||||
session_id: str,
|
session_id: str,
|
||||||
user_id: str,
|
user_id: str,
|
||||||
message: str,
|
message: str,
|
||||||
channel: str = None,
|
images: List[str] = None,
|
||||||
source: str = None,
|
channel: str = None,
|
||||||
username: str = None,
|
source: str = None,
|
||||||
|
username: str = None,
|
||||||
) -> str:
|
) -> str:
|
||||||
"""
|
"""
|
||||||
处理用户消息:将消息放入会话队列,按顺序依次处理。
|
处理用户消息:将消息放入会话队列,按顺序依次处理。
|
||||||
@@ -483,6 +497,7 @@ class AgentManager:
|
|||||||
session_id=session_id,
|
session_id=session_id,
|
||||||
user_id=user_id,
|
user_id=user_id,
|
||||||
message=message,
|
message=message,
|
||||||
|
images=images,
|
||||||
channel=channel,
|
channel=channel,
|
||||||
source=source,
|
source=source,
|
||||||
username=username,
|
username=username,
|
||||||
@@ -497,8 +512,8 @@ class AgentManager:
|
|||||||
|
|
||||||
# 如果队列中已有等待的消息,通知用户消息已排队
|
# 如果队列中已有等待的消息,通知用户消息已排队
|
||||||
if queue_size > 0 or (
|
if queue_size > 0 or (
|
||||||
session_id in self._session_workers
|
session_id in self._session_workers
|
||||||
and not self._session_workers[session_id].done()
|
and not self._session_workers[session_id].done()
|
||||||
):
|
):
|
||||||
logger.info(
|
logger.info(
|
||||||
f"会话 {session_id} 有任务正在处理,消息已排队等待 "
|
f"会话 {session_id} 有任务正在处理,消息已排队等待 "
|
||||||
@@ -510,8 +525,8 @@ class AgentManager:
|
|||||||
|
|
||||||
# 确保该会话有一个worker在运行
|
# 确保该会话有一个worker在运行
|
||||||
if (
|
if (
|
||||||
session_id not in self._session_workers
|
session_id not in self._session_workers
|
||||||
or self._session_workers[session_id].done()
|
or self._session_workers[session_id].done()
|
||||||
):
|
):
|
||||||
self._session_workers[session_id] = asyncio.create_task(
|
self._session_workers[session_id] = asyncio.create_task(
|
||||||
self._session_worker(session_id)
|
self._session_worker(session_id)
|
||||||
@@ -552,8 +567,8 @@ class AgentManager:
|
|||||||
self._session_workers.pop(session_id, None) # noqa
|
self._session_workers.pop(session_id, None) # noqa
|
||||||
# 如果队列为空,清理队列
|
# 如果队列为空,清理队列
|
||||||
if (
|
if (
|
||||||
session_id in self._session_queues
|
session_id in self._session_queues
|
||||||
and self._session_queues[session_id].empty()
|
and self._session_queues[session_id].empty()
|
||||||
):
|
):
|
||||||
self._session_queues.pop(session_id, None)
|
self._session_queues.pop(session_id, None)
|
||||||
|
|
||||||
@@ -584,7 +599,7 @@ class AgentManager:
|
|||||||
if task.username:
|
if task.username:
|
||||||
agent.username = task.username
|
agent.username = task.username
|
||||||
|
|
||||||
return await agent.process(task.message)
|
return await agent.process(task.message, images=task.images)
|
||||||
|
|
||||||
async def clear_session(self, session_id: str, user_id: str):
|
async def clear_session(self, session_id: str, user_id: str):
|
||||||
"""
|
"""
|
||||||
|
|||||||
+608
-301
File diff suppressed because it is too large
Load Diff
@@ -130,10 +130,11 @@ class DiscordModule(_ModuleBase, _MessageBase[Discord]):
|
|||||||
if msg_type == "message":
|
if msg_type == "message":
|
||||||
text = msg_json.get("text")
|
text = msg_json.get("text")
|
||||||
chat_id = msg_json.get("chat_id")
|
chat_id = msg_json.get("chat_id")
|
||||||
if text and userid:
|
images = self._extract_images(msg_json)
|
||||||
|
if (text or images) and userid:
|
||||||
logger.info(
|
logger.info(
|
||||||
f"收到来自 {client_config.name} 的 Discord 消息:"
|
f"收到来自 {client_config.name} 的 Discord 消息:"
|
||||||
f"userid={userid}, username={username}, text={text}"
|
f"userid={userid}, username={username}, text={text}, images={len(images) if images else 0}"
|
||||||
)
|
)
|
||||||
return CommingMessage(
|
return CommingMessage(
|
||||||
channel=MessageChannel.Discord,
|
channel=MessageChannel.Discord,
|
||||||
@@ -142,9 +143,26 @@ class DiscordModule(_ModuleBase, _MessageBase[Discord]):
|
|||||||
username=username,
|
username=username,
|
||||||
text=text,
|
text=text,
|
||||||
chat_id=str(chat_id) if chat_id else None,
|
chat_id=str(chat_id) if chat_id else None,
|
||||||
|
images=images,
|
||||||
)
|
)
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _extract_images(msg_json: dict) -> Optional[List[str]]:
|
||||||
|
"""
|
||||||
|
从Discord消息中提取图片URL
|
||||||
|
"""
|
||||||
|
attachments = msg_json.get("attachments", [])
|
||||||
|
if not attachments:
|
||||||
|
return None
|
||||||
|
images = []
|
||||||
|
for attachment in attachments:
|
||||||
|
if attachment.get("type") == "image":
|
||||||
|
url = attachment.get("url")
|
||||||
|
if url:
|
||||||
|
images.append(url)
|
||||||
|
return images if images else None
|
||||||
|
|
||||||
def post_message(self, message: Notification, **kwargs) -> None:
|
def post_message(self, message: Notification, **kwargs) -> None:
|
||||||
"""
|
"""
|
||||||
发送通知消息
|
发送通知消息
|
||||||
|
|||||||
@@ -198,10 +198,12 @@ class SlackModule(_ModuleBase, _MessageBase[Slack]):
|
|||||||
logger.debug(f"解析Slack消息失败:{str(err)}")
|
logger.debug(f"解析Slack消息失败:{str(err)}")
|
||||||
return None
|
return None
|
||||||
if msg_json:
|
if msg_json:
|
||||||
|
images = None
|
||||||
if msg_json.get("type") == "message":
|
if msg_json.get("type") == "message":
|
||||||
userid = msg_json.get("user")
|
userid = msg_json.get("user")
|
||||||
text = msg_json.get("text")
|
text = msg_json.get("text")
|
||||||
username = msg_json.get("user")
|
username = msg_json.get("user")
|
||||||
|
images = self._extract_images(msg_json)
|
||||||
elif msg_json.get("type") == "block_actions":
|
elif msg_json.get("type") == "block_actions":
|
||||||
userid = msg_json.get("user", {}).get("id")
|
userid = msg_json.get("user", {}).get("id")
|
||||||
callback_data = msg_json.get("actions")[0].get("value")
|
callback_data = msg_json.get("actions")[0].get("value")
|
||||||
@@ -243,6 +245,7 @@ class SlackModule(_ModuleBase, _MessageBase[Slack]):
|
|||||||
flags=re.IGNORECASE,
|
flags=re.IGNORECASE,
|
||||||
).strip()
|
).strip()
|
||||||
username = ""
|
username = ""
|
||||||
|
images = self._extract_images(msg_json.get("event", {}))
|
||||||
elif msg_json.get("type") == "shortcut":
|
elif msg_json.get("type") == "shortcut":
|
||||||
userid = msg_json.get("user", {}).get("id")
|
userid = msg_json.get("user", {}).get("id")
|
||||||
text = msg_json.get("callback_id")
|
text = msg_json.get("callback_id")
|
||||||
@@ -254,7 +257,7 @@ class SlackModule(_ModuleBase, _MessageBase[Slack]):
|
|||||||
else:
|
else:
|
||||||
return None
|
return None
|
||||||
logger.info(
|
logger.info(
|
||||||
f"收到来自 {client_config.name} 的Slack消息:userid={userid}, username={username}, text={text}"
|
f"收到来自 {client_config.name} 的Slack消息:userid={userid}, username={username}, text={text}, images={len(images) if images else 0}"
|
||||||
)
|
)
|
||||||
return CommingMessage(
|
return CommingMessage(
|
||||||
channel=MessageChannel.Slack,
|
channel=MessageChannel.Slack,
|
||||||
@@ -262,9 +265,26 @@ class SlackModule(_ModuleBase, _MessageBase[Slack]):
|
|||||||
userid=userid,
|
userid=userid,
|
||||||
username=username,
|
username=username,
|
||||||
text=text,
|
text=text,
|
||||||
|
images=images,
|
||||||
)
|
)
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _extract_images(msg_json: dict) -> Optional[List[str]]:
|
||||||
|
"""
|
||||||
|
从Slack消息中提取图片URL
|
||||||
|
"""
|
||||||
|
files = msg_json.get("files", [])
|
||||||
|
if not files:
|
||||||
|
return None
|
||||||
|
images = []
|
||||||
|
for file in files:
|
||||||
|
if file.get("type") in ("image", "jpg", "jpeg", "png", "gif", "webp"):
|
||||||
|
url = file.get("url_private") or file.get("url_private_download")
|
||||||
|
if url:
|
||||||
|
images.append(url)
|
||||||
|
return images if images else None
|
||||||
|
|
||||||
def post_message(self, message: Notification, **kwargs) -> None:
|
def post_message(self, message: Notification, **kwargs) -> None:
|
||||||
"""
|
"""
|
||||||
发送消息
|
发送消息
|
||||||
|
|||||||
@@ -194,26 +194,33 @@ class TelegramModule(_ModuleBase, _MessageBase[Telegram]):
|
|||||||
text = msg.get("text")
|
text = msg.get("text")
|
||||||
user_id = msg.get("from", {}).get("id")
|
user_id = msg.get("from", {}).get("id")
|
||||||
user_name = msg.get("from", {}).get("username")
|
user_name = msg.get("from", {}).get("username")
|
||||||
# Extract chat_id to enable correct reply targeting
|
|
||||||
chat_id = msg.get("chat", {}).get("id")
|
chat_id = msg.get("chat", {}).get("id")
|
||||||
|
|
||||||
if text and user_id:
|
images = self._extract_images(msg)
|
||||||
|
|
||||||
|
if user_id:
|
||||||
|
if not text and not images:
|
||||||
|
logger.debug(
|
||||||
|
f"收到来自 {client_config.name} 的Telegram消息无文本和图片"
|
||||||
|
)
|
||||||
|
return None
|
||||||
|
|
||||||
logger.info(
|
logger.info(
|
||||||
f"收到来自 {client_config.name} 的Telegram消息:"
|
f"收到来自 {client_config.name} 的Telegram消息:"
|
||||||
f"userid={user_id}, username={user_name}, chat_id={chat_id}, text={text}"
|
f"userid={user_id}, username={user_name}, chat_id={chat_id}, text={text}, images={len(images) if images else 0}"
|
||||||
)
|
)
|
||||||
|
|
||||||
# Clean bot mentions from text to ensure consistent processing
|
cleaned_text = (
|
||||||
cleaned_text = self._clean_bot_mention(
|
self._clean_bot_mention(text, client.bot_username if client else None)
|
||||||
text, client.bot_username if client else None
|
if text
|
||||||
|
else None
|
||||||
)
|
)
|
||||||
|
|
||||||
# 检查权限
|
|
||||||
admin_users = client_config.config.get("TELEGRAM_ADMINS")
|
admin_users = client_config.config.get("TELEGRAM_ADMINS")
|
||||||
user_list = client_config.config.get("TELEGRAM_USERS")
|
user_list = client_config.config.get("TELEGRAM_USERS")
|
||||||
config_chat_id = client_config.config.get("TELEGRAM_CHAT_ID")
|
config_chat_id = client_config.config.get("TELEGRAM_CHAT_ID")
|
||||||
|
|
||||||
if cleaned_text.startswith("/"):
|
if cleaned_text and cleaned_text.startswith("/"):
|
||||||
if (
|
if (
|
||||||
admin_users
|
admin_users
|
||||||
and str(user_id) not in admin_users.split(",")
|
and str(user_id) not in admin_users.split(",")
|
||||||
@@ -236,11 +243,34 @@ class TelegramModule(_ModuleBase, _MessageBase[Telegram]):
|
|||||||
source=client_config.name,
|
source=client_config.name,
|
||||||
userid=user_id,
|
userid=user_id,
|
||||||
username=user_name,
|
username=user_name,
|
||||||
text=cleaned_text, # Use cleaned text
|
text=cleaned_text,
|
||||||
chat_id=str(chat_id) if chat_id else None,
|
chat_id=str(chat_id) if chat_id else None,
|
||||||
|
images=images if images else None,
|
||||||
)
|
)
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _extract_images(msg: dict) -> Optional[List[str]]:
|
||||||
|
"""
|
||||||
|
从Telegram消息中提取图片file_id
|
||||||
|
"""
|
||||||
|
images = []
|
||||||
|
photo = msg.get("photo")
|
||||||
|
if photo and isinstance(photo, list):
|
||||||
|
largest_photo = photo[-1]
|
||||||
|
file_id = largest_photo.get("file_id")
|
||||||
|
if file_id:
|
||||||
|
images.append(file_id)
|
||||||
|
|
||||||
|
document = msg.get("document")
|
||||||
|
if document:
|
||||||
|
file_id = document.get("file_id")
|
||||||
|
mime_type = document.get("mime_type", "")
|
||||||
|
if file_id and mime_type.startswith("image/"):
|
||||||
|
images.append(file_id)
|
||||||
|
|
||||||
|
return images if images else None
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _clean_bot_mention(text: str, bot_username: Optional[str]) -> str:
|
def _clean_bot_mention(text: str, bot_username: Optional[str]) -> str:
|
||||||
"""
|
"""
|
||||||
|
|||||||
@@ -53,6 +53,8 @@ class CommingMessage(BaseModel):
|
|||||||
chat_id: Optional[str] = None
|
chat_id: Optional[str] = None
|
||||||
# 完整的回调查询信息(原始数据)
|
# 完整的回调查询信息(原始数据)
|
||||||
callback_query: Optional[Dict] = None
|
callback_query: Optional[Dict] = None
|
||||||
|
# 图片列表(图片URL或file_id)
|
||||||
|
images: Optional[List[str]] = None
|
||||||
|
|
||||||
def to_dict(self):
|
def to_dict(self):
|
||||||
"""
|
"""
|
||||||
|
|||||||
Reference in New Issue
Block a user