diff --git a/README.md b/README.md index 4a90746..c122b60 100644 --- a/README.md +++ b/README.md @@ -120,6 +120,9 @@ _✨ 支持多API预设、MCP协议、内置工具、联网搜索、视觉模型 | LLMCHAT__HISTORY_SIZE | 否 | 20 | LLM上下文消息保留数量(1-40),越大token消耗量越多 | | LLMCHAT__PAST_EVENTS_SIZE | 否 | 10 | 触发回复时发送的群消息数量(1-20),越大token消耗量越多 | | LLMCHAT__REQUEST_TIMEOUT | 否 | 30 | API请求超时时间(秒) | +| LLMCHAT__MAX_TOOL_ROUNDS | 否 | 8 | 单次对话最大工具调用轮数,达到后强制生成最终回答 | +| LLMCHAT__MAX_REPEATED_TOOL_CALLS | 否 | 2 | 工具名及参数完全相同时允许实际执行的最大次数 | +| LLMCHAT__MCP_TIMEOUT | 否 | 30 | MCP建连、工具发现及执行的超时时间(秒) | | LLMCHAT__DEFAULT_PRESET | 否 | off | 默认使用的预设名称,配置为off则为关闭 | | LLMCHAT__RANDOM_TRIGGER_PROB | 否 | 0.05 | 默认随机触发概率 [0, 1] | | LLMCHAT__DEFAULT_PROMPT | 否 | 你的回答应该尽量简洁、幽默、可以使用一些语气词、颜文字。你应该拒绝回答任何政治相关的问题。 | 默认提示词 | @@ -267,6 +270,23 @@ LLMCHAT__MCP_SERVERS同样为一个dict,key为服务器名称,value配置的 +## 🧩 代码结构 + +核心代码按职责拆分: + +- `__init__.py`:仅负责 NoneBot 注册、依赖装配和生命周期。 +- `state.py`:群聊/私聊状态与预设选择。 +- `dispatcher.py`:每会话单 worker 的可靠队列调度。 +- `conversation.py`:LLM 请求、有限工具循环和回复发送。 +- `message_utils.py` 与 `prompts.py`:消息转换、图片处理和提示词构建。 +- `mcpclient.py` 与 `onebottools.py`:外部 MCP 和 OneBot 工具适配。 +- `persistence.py`:带锁和原子替换的状态持久化。 +- `commands.py`:管理命令注册。 + +运行回归测试: + + python -m unittest discover -s tests -v + ## 🎉 使用 **如果`LLMCHAT__DEFAULT_PRESET`没有配置,则插件默认为关闭状态,请使用`API预设+[预设名]`开启插件, 私聊同理。** diff --git a/nonebot_plugin_llmchat/__init__.py b/nonebot_plugin_llmchat/__init__.py index af6f315..d27a497 100755 --- a/nonebot_plugin_llmchat/__init__.py +++ b/nonebot_plugin_llmchat/__init__.py @@ -1,36 +1,17 @@ -import asyncio -import base64 -from collections import defaultdict, deque -from datetime import datetime -import json -import os import random -import re -import ssl -import time -from typing import TYPE_CHECKING -import aiofiles -import httpx -from nonebot import ( - get_bot, - get_driver, - get_plugin_config, - logger, - on_command, - on_message, - require, -) -from nonebot.adapters.onebot.v11 import GroupMessageEvent, Message, MessageSegment, PrivateMessageEvent -from nonebot.adapters.onebot.v11.permission import GROUP_ADMIN, GROUP_OWNER, PRIVATE -from nonebot.params import CommandArg -from nonebot.permission import SUPERUSER +from nonebot import get_driver, get_plugin_config, logger, on_message, require +from nonebot.adapters.onebot.v11 import GroupMessageEvent, PrivateMessageEvent from nonebot.plugin import PluginMetadata from nonebot.rule import Rule -from openai import AsyncOpenAI -from .config import Config, PresetConfig +from .commands import register_commands +from .config import Config +from .conversation import ConversationService +from .dispatcher import MessageDispatcher from .mcpclient import MCPClient +from .persistence import StatePersistence +from .state import StateStore require("nonebot_plugin_localstore") import nonebot_plugin_localstore as store @@ -38,16 +19,10 @@ import nonebot_plugin_localstore as store require("nonebot_plugin_apscheduler") from nonebot_plugin_apscheduler import scheduler -if TYPE_CHECKING: - from openai.types.chat import ( - ChatCompletionContentPartParam, - ChatCompletionMessageParam, - ) - __plugin_meta__ = PluginMetadata( name="llmchat", - description="支持多API预设、MCP协议、联网搜索、视觉模型、Nano Banana(生图模型)的AI群聊插件", - usage="""@机器人 + 消息 开启对话""", + description="支持多API预设、MCP协议、联网搜索、视觉模型的AI群聊插件", + usage="@机器人 + 消息 开启对话", type="application", homepage="https://github.com/FuQuan233/nonebot-plugin-llmchat", config=Config, @@ -56,852 +31,76 @@ __plugin_meta__ = PluginMetadata( plugin_config = get_plugin_config(Config).llmchat driver = get_driver() -tasks: set["asyncio.Task"] = set() +states = StateStore(plugin_config) - -def pop_reasoning_content( - content: str | None, -) -> tuple[str | None, str | None]: - if content is None: - return None, None - - # 如果找到了 标签内容,返回过滤后的文本和标签内的内容,否则只返回过滤后的文本和None - if matched := re.match(r"(.*?)", content, flags=re.DOTALL): - reasoning_element = matched.group(0) - reasoning_content = matched.group(1).strip() - filtered_content = content.replace(reasoning_element, "").strip() - - return filtered_content, reasoning_content - else: - return content, None - - -# 初始化群组状态 -class GroupState: - def __init__(self): - self.preset_name = plugin_config.default_preset - self.history = deque(maxlen=plugin_config.history_size * 2) - self.queue = asyncio.Queue() - self.processing = False - self.last_active = time.time() - self.past_events = deque(maxlen=plugin_config.past_events_size) - self.group_prompt: str | None = None - self.user_prompt: str | None = None - self.output_reasoning_content = False - self.random_trigger_prob = plugin_config.random_trigger_prob - - -# 初始化私聊状态 -class PrivateChatState: - def __init__(self): - self.preset_name = plugin_config.private_chat_preset - self.history = deque(maxlen=plugin_config.history_size * 2) - self.queue = asyncio.Queue() - self.processing = False - self.last_active = time.time() - self.past_events = deque(maxlen=plugin_config.past_events_size) - self.group_prompt: str | None = None - self.output_reasoning_content = False - - -group_states: dict[int, GroupState] = defaultdict(GroupState) -private_chat_states: dict[int, PrivateChatState] = defaultdict(PrivateChatState) - - -# 获取当前预设配置 -def get_preset(context_id: int, is_group: bool = True) -> PresetConfig: - if is_group: - state = group_states[context_id] - else: - state = private_chat_states[context_id] - - for preset in plugin_config.api_presets: - if preset.name == state.preset_name: - return preset - return plugin_config.api_presets[0] # 默认返回第一个预设 - - -# 消息格式转换 -def format_message(event: GroupMessageEvent | PrivateMessageEvent) -> str: - text_message = "" - if isinstance(event, GroupMessageEvent) and event.reply is not None: - text_message += f"[回复 {event.reply.sender.nickname} 的消息 {event.reply.message.extract_plain_text()}]\n" - - if isinstance(event, GroupMessageEvent) and event.is_tome(): - text_message += f"@{next(iter(driver.config.nickname))} " - - for msgseg in event.get_message(): - if msgseg.type == "at": - text_message += msgseg.data.get("name", "") - elif msgseg.type == "image": - text_message += "[图片]" - elif msgseg.type == "voice": - text_message += "[语音]" - elif msgseg.type == "face": - pass - elif msgseg.type == "text": - text_message += msgseg.data.get("text", "") - - if isinstance(event, GroupMessageEvent): - message = { - "SenderNickname": str(event.sender.card or event.sender.nickname), - "SenderUserId": str(event.user_id), - "Message": text_message, - "MessageID": event.message_id, - "SendTime": datetime.fromtimestamp(event.time).isoformat(), - } - else: # PrivateMessageEvent - message = { - "SenderNickname": str(event.sender.nickname), - "SenderUserId": str(event.user_id), - "Message": text_message, - "MessageID": event.message_id, - "SendTime": datetime.fromtimestamp(event.time).isoformat(), - } - return json.dumps(message, ensure_ascii=False) - - -def build_reasoning_forward_nodes(self_id: str, reasoning_content: str): - self_nickname = next(iter(driver.config.nickname)) - nodes = [ - { - "type": "node", - "data": { - "nickname": self_nickname, - "user_id": self_id, - "content": f"{self_nickname}的内心OS:", - }, - }, - { - "type": "node", - "data": { - "nickname": self_nickname, - "user_id": self_id, - "content": reasoning_content, - }, - }, - ] - - return nodes +# 保留旧版公开名称,避免依赖插件内部状态的代码立即失效。 +group_states = states.group_states +private_chat_states = states.private_states async def is_triggered(event: GroupMessageEvent | PrivateMessageEvent) -> bool: - """扩展后的消息处理规则""" - - if isinstance(event, GroupMessageEvent): - state = group_states[event.group_id] - - if state.preset_name == "off": - return False - - # 黑名单用户 - if event.user_id in plugin_config.blacklist_user_ids: - return False - - # 忽略特定前缀的消息 - msg_text = event.get_plaintext().strip() - for prefix in plugin_config.ignore_prefixes: - if msg_text.startswith(prefix): - return False - - state.past_events.append(event) - - # 原有@触发条件 - if event.is_tome(): - return True - - # 随机触发条件 - if random.random() < state.random_trigger_prob: - return True - + is_group = isinstance(event, GroupMessageEvent) + if not is_group and not plugin_config.enable_private_chat: + return False + state = states.get(event.group_id if is_group else event.user_id, is_group) + if state.preset_name == "off" or event.user_id in plugin_config.blacklist_user_ids: + return False + text = event.get_plaintext().strip() + if any(text.startswith(prefix) for prefix in plugin_config.ignore_prefixes): return False - elif isinstance(event, PrivateMessageEvent): - # 检查私聊功能是否启用 - if not plugin_config.enable_private_chat: - return False - - state = private_chat_states[event.user_id] - - if state.preset_name == "off": - return False - - # 黑名单用户 - if event.user_id in plugin_config.blacklist_user_ids: - return False - - # 忽略特定前缀的消息 - msg_text = event.get_plaintext().strip() - for prefix in plugin_config.ignore_prefixes: - if msg_text.startswith(prefix): - return False - - state.past_events.append(event) - - # 私聊默认触发 + state.pending_events.append(event) + if not is_group: return True - - return False + return event.is_tome() or random.random() < state.random_trigger_prob -# 消息处理器 -handler = on_message( - rule=Rule(is_triggered), - priority=99, - block=False, +handler = on_message(rule=Rule(is_triggered), priority=99, block=False) +conversation = ConversationService( + config=plugin_config, + states=states, + bot_names=set(driver.config.nickname), + sender=handler.send, ) +message_dispatcher = MessageDispatcher(conversation.process_event) @handler.handle() -async def handle_message(event: GroupMessageEvent | PrivateMessageEvent): - if isinstance(event, GroupMessageEvent): - group_id = event.group_id - logger.debug( - f"收到群聊消息 群号:{group_id} 用户:{event.user_id} 内容:{event.get_plaintext()}" - ) - state = group_states[group_id] - context_id = group_id - else: # PrivateMessageEvent - user_id = event.user_id - logger.debug( - f"收到私聊消息 用户:{user_id} 内容:{event.get_plaintext()}" - ) - state = private_chat_states[user_id] - context_id = user_id - - await state.queue.put(event) - if not state.processing: - state.processing = True - is_group = isinstance(event, GroupMessageEvent) - task = asyncio.create_task(process_messages(context_id, is_group)) - task.add_done_callback(tasks.discard) - tasks.add(task) - -async def process_images(event: GroupMessageEvent | PrivateMessageEvent) -> list[str]: - base64_images = [] - for segement in event.get_message(): - if segement.type == "image": - image_url = segement.data.get("url") or segement.data.get("file") - if image_url: - try: - # 处理高版本 httpx 的 [SSL: SSLV3_ALERT_HANDSHAKE_FAILURE] 报错 - ssl_context = ssl.create_default_context() - ssl_context.check_hostname = False - ssl_context.verify_mode = ssl.CERT_NONE - ssl_context.set_ciphers("DEFAULT@SECLEVEL=2") - - # 下载图片并将图片转换为base64 - async with httpx.AsyncClient(verify=ssl_context) as client: - response = await client.get(image_url, timeout=10.0) - if response.status_code != 200: - logger.error(f"下载图片失败: {image_url}, 状态码: {response.status_code}") - continue - image_data = response.content - base64_data = base64.b64encode(image_data).decode("utf-8") - base64_images.append(base64_data) - except Exception as e: - logger.error(f"处理图片时出错: {e}") - logger.debug(f"共处理 {len(base64_images)} 张图片") - return base64_images - -async def send_split_messages(message_handler, content: str): - """ - 将消息按分隔符分段并发送 - """ - logger.info(f"准备发送分段消息,分段数:{len(content.split(''))}") - for segment in content.split(""): - # 跳过空消息 - if not segment.strip(): - continue - segment = segment.strip() # 删除前后多余的换行和空格 - await asyncio.sleep(2) # 避免发送过快 - logger.debug(f"发送消息分段 内容:{segment[:50]}...") # 只记录前50个字符避免日志过大 - await message_handler.send(Message(segment)) - -async def process_messages(context_id: int, is_group: bool = True): - if is_group: - group_id = context_id - state = group_states[group_id] - else: - user_id = context_id - state = private_chat_states[user_id] - group_id = None - - preset = get_preset(context_id, is_group) - - # 初始化OpenAI客户端 - if preset.proxy != "": - client = AsyncOpenAI( - base_url=preset.api_base, - api_key=preset.api_key, - timeout=plugin_config.request_timeout, - http_client=httpx.AsyncClient(proxy=preset.proxy), - ) - else: - client = AsyncOpenAI( - base_url=preset.api_base, - api_key=preset.api_key, - timeout=plugin_config.request_timeout, - ) - - chat_type = "群聊" if is_group else "私聊" - context_type = "群号" if is_group else "用户" - logger.info( - f"开始处理{chat_type}消息 {context_type}:{context_id} 当前队列长度:{state.queue.qsize()}" +async def handle_message(event: GroupMessageEvent | PrivateMessageEvent) -> None: + is_group = isinstance(event, GroupMessageEvent) + context_id = event.group_id if is_group else event.user_id + state = states.get(context_id, is_group) + logger.debug( + f"收到{'群聊' if is_group else '私聊'}消息 " f"会话:{context_id} 用户:{event.user_id} 内容:{event.get_plaintext()}" ) - try: - while not state.queue.empty(): - event = await state.queue.get() - if is_group: - logger.debug(f"从队列获取消息 群号:{context_id} 消息ID:{event.message_id}") - group_id = context_id - else: - logger.debug(f"从队列获取消息 用户:{context_id} 消息ID:{event.message_id}") - group_id = None - past_events_snapshot = [] - mcp_client = MCPClient.get_instance( - plugin_config.mcp_servers, - plugin_config.mcp_server_cwd, - ) - try: - # 构建系统提示,分成多行以满足行长限制 - chat_type = "群聊" if is_group else "私聊" - bot_names = "、".join(list(driver.config.nickname)) - default_prompt = (state.group_prompt) or plugin_config.default_prompt - - system_lines = [ - f"我想要你帮我在{chat_type}中闲聊,大家一般叫你{bot_names}。", - "我将会在后面的信息中告诉你每条信息的发送者和发送时间,你可以直接称呼发送者为他对应的昵称。", - "你的回复需要遵守以下几点规则:", - "- 你可以使用多条消息回复,每两条消息之间使用分隔,前后不需要包含额外的换行和空格。", - "- 除外,消息中不应该包含其他类似的标记。", - "- 不要使用markdown或者html,聊天软件不支持解析,换行请用换行符。", - "- 你应该以普通人的方式发送消息,每条消息字数要尽量少一些,应该倾向于使用更多条的消息回复。", - "- 代码则不需要分段,用单独的一条消息发送。", - "- 请使用发送者的昵称称呼发送者,你可以礼貌地问候发送者,但只需要在" - "第一次回答这位发送者的问题时问候他。", - "- 你有引用某条消息的能力,使用[CQ:reply,id=(消息id)]来引用。", - "- 如果有多条消息,你应该优先回复提到你的,一段时间之前的就不要回复了,也可以直接选择不回复。", - "- 如果你选择完全不回复,你只需要直接输出一个。", - "- 如果你需要思考的话,你应该尽量少思考,以节省时间。", - ] - - if is_group: - system_lines += [ - "- 你有at群成员的能力,只需要在某条消息中插入[CQ:at,qq=(QQ号)]," - "也就是CQ码。at发送者是非必要的,你可以根据你自己的想法at某个人。", - ] - - system_lines += [ - "下面是关于你性格的设定,如果设定中提到让你扮演某个人,或者设定中有提到名字,则优先使用设定中的名字。", - default_prompt, - ] - - systemPrompt = "\n".join(system_lines) - if preset.support_mcp: - systemPrompt += "\n你也可以使用一些工具,下面是关于这些工具的额外说明:\n" - for mcp_name, mcp_config in plugin_config.mcp_servers.items(): - if mcp_config.additional_prompt: - systemPrompt += f"{mcp_name}:{mcp_config.additional_prompt}" - systemPrompt += "\n" - - logger.debug(f"构建系统提示词:\n{systemPrompt}") - - messages: list[ChatCompletionMessageParam] = [ - {"role": "system", "content": systemPrompt} - ] - - while len(state.history) > 0 and state.history[0]["role"] != "user": - state.history.popleft() - - messages += list(state.history)[-plugin_config.history_size * 2 :] - - # 没有未处理的消息说明已经被处理了,跳过 - if state.past_events.__len__() < 1: - break - - content: list[ChatCompletionContentPartParam] = [] - - # 将机器人错过的消息推送给LLM - past_events_snapshot = list(state.past_events) - state.past_events.clear() - for ev in past_events_snapshot: - text_content = format_message(ev) - content.append({"type": "text", "text": text_content}) - - # 将消息中的图片转成 base64 - if preset.support_image: - base64_images = await process_images(ev) - for base64_image in base64_images: - content.append({"type": "image_url", "image_url": {"url": f"data:image/jpeg;base64,{base64_image}"}}) - - new_messages: list[ChatCompletionMessageParam] = [ - {"role": "user", "content": content} - ] - - logger.debug( - f"发送API请求 模型:{preset.model_name} 历史消息数:{len(messages)}" - ) - - client_config = { - "model": preset.model_name, - "max_tokens": preset.max_tokens, - "temperature": preset.temperature, - "timeout": 60, - "extra_body": preset.extra_body, - } - - if preset.support_mcp: - available_tools = await mcp_client.get_available_tools(is_group) - client_config["tools"] = available_tools - - response = await client.chat.completions.create( - **client_config, - messages=messages + new_messages, - ) - - if response.usage is not None: - logger.debug(f"收到API响应 使用token数:{response.usage.total_tokens}") - - message = response.choices[0].message - - # 处理响应并处理工具调用 - while preset.support_mcp and message and message.tool_calls: - llm_reply: ChatCompletionMessageParam = { - "role": "assistant", - "content": message.content, - "tool_calls": [tool_call.model_dump() for tool_call in message.tool_calls] - } - - if preset.request_with_reasoning_content: - llm_reply["reasoning_content"] = message.reasoning_content # pyright: ignore[reportGeneralTypeIssues] - - # 发送LLM调用工具时的回复,一般没有 - if message.content: - await send_split_messages(handler, message.content) - - # 处理每个工具调用 - new_messages.append(llm_reply) - - for tool_call in message.tool_calls: - logger.debug(f"处理工具调用:{tool_call.function.name} 参数:{tool_call.function.arguments}") - - tool_name = tool_call.function.name - try: - tool_args = json.loads(tool_call.function.arguments) - except (json.JSONDecodeError, TypeError, ValueError) as e: - error_message = ( - f"工具调用参数格式错误,无法解析 {tool_name} 的 arguments: {e!s}. " - f"原始参数: {tool_call.function.arguments}" - ) - logger.warning(error_message) - new_messages.append({ - "role": "tool", - "tool_call_id": tool_call.id, - "content": error_message, - }) - continue - - # 发送工具调用提示 - await handler.send(Message(f"正在使用{mcp_client.get_friendly_name(tool_name)}")) - - if is_group: - result = await mcp_client.call_tool( - tool_name, - tool_args, - group_id=event.group_id, - bot_id=str(event.self_id) - ) - else: - result = await mcp_client.call_tool( - tool_name, - tool_args, - bot_id=str(event.self_id) - ) - - new_messages.append({ - "role": "tool", - "tool_call_id": tool_call.id, - "content": str(result) - }) - - # 将工具调用的结果交给 LLM - response = await client.chat.completions.create( - **client_config, - messages=messages + new_messages, - ) - - message = response.choices[0].message - - # 安全检查:确保 message 不为 None - if not message: - logger.error("API 响应中的 message 为 None") - await handler.send(Message("服务暂时不可用,请稍后再试")) - return - - reply, matched_reasoning_content = pop_reasoning_content( - message.content - ) - reasoning_content: str | None = ( - getattr(message, "reasoning_content", None) - or matched_reasoning_content - ) - - llm_reply: ChatCompletionMessageParam = { - "role": "assistant", - "content": reply, - } - - reply_images = getattr(message, "images", None) - - if reply_images: - # openai的sdk里的assistant消息暂时没有images字段,需要单独处理 - llm_reply["images"] = reply_images # pyright: ignore[reportGeneralTypeIssues] - - if preset.request_with_reasoning_content: - llm_reply["reasoning_content"] = reasoning_content # pyright: ignore[reportGeneralTypeIssues] - - new_messages.append(llm_reply) - - # 请求成功后再保存历史记录,保证user和assistant穿插,防止R1模型报错 - for message in new_messages: - state.history.append(message) - - if state.output_reasoning_content and reasoning_content: - try: - bot = get_bot(str(event.self_id)) - if is_group: - await bot.send_group_forward_msg( - group_id=group_id, - messages=build_reasoning_forward_nodes( - bot.self_id, reasoning_content - ), - ) - else: - await bot.send_private_forward_msg( - user_id=context_id, - messages=build_reasoning_forward_nodes( - bot.self_id, reasoning_content - ), - ) - except Exception as e: - logger.error(f"合并转发消息发送失败:\n{e!s}\n") - - assert reply is not None - await send_split_messages(handler, reply) - - if reply_images: - logger.debug(f"API响应 图片数:{len(reply_images)}") - for i, image in enumerate(reply_images, start=1): - logger.debug(f"正在发送第{i}张图片") - image_base64 = image["image_url"]["url"].removeprefix("data:image/png;base64,") - image_msg = MessageSegment.image(base64.b64decode(image_base64)) - await handler.send(image_msg) - - except Exception as e: - logger.opt(exception=e).error(f"API请求失败 {'群号' if is_group else '用户'}:{context_id}") - # 如果在处理过程中出现异常,恢复未处理的消息到state中 - state.past_events.extendleft(reversed(past_events_snapshot)) - await handler.send(Message(f"服务暂时不可用,请稍后再试\n{e!s}")) - finally: - state.queue.task_done() - # 不再需要每次都清理MCPClient,因为它现在是单例 - # await mcp_client.cleanup() - finally: - state.processing = False + await message_dispatcher.enqueue(context_id, is_group, state, event) -# 预设切换命令 -preset_handler = on_command("API预设", priority=1, block=True, permission=SUPERUSER) - - -@preset_handler.handle() -async def handle_preset(event: GroupMessageEvent | PrivateMessageEvent, args: Message = CommandArg()): - # 解析命令参数 - args_text = args.extract_plain_text().strip() - args_parts = args_text.split(maxsplit=1) - - target_id = None - preset_name = None - - # 可用预设列表 - available_presets = {p.name for p in plugin_config.api_presets} - - # 只在私聊中允许 SUPERUSER 修改他人预设 - if isinstance(event, PrivateMessageEvent) and args_parts and args_parts[0].isdigit(): - # 第一个参数是纯数字,且不是预设名 - if args_parts[0] not in available_presets: - target_id = int(args_parts[0]) - - # 判断目标是群聊还是私聊 - if target_id in group_states: - state = group_states[target_id] - is_group_target = True - elif target_id in private_chat_states: - state = private_chat_states[target_id] - is_group_target = False - else: - # 默认创建私聊状态 - state = private_chat_states[target_id] - is_group_target = False - - # 如果只有目标 ID,没有预设名,返回当前预设 - if len(args_parts) == 1: - context_type = "群聊" if is_group_target else "私聊" - available_presets_str = "\n- ".join(available_presets) - await preset_handler.finish( - f"{context_type} {target_id} 当前API预设:{state.preset_name}\n可用API预设:\n- {available_presets_str}" - ) - - # 有预设名,进行修改 - preset_name = args_parts[1] - context_id = target_id - else: - # 第一个参数虽然是数字但也是预设名,按普通流程处理 - target_id = None - preset_name = args_text - if not plugin_config.enable_private_chat: - return - context_id = event.user_id - state = private_chat_states[context_id] - is_group_target = False - else: - # 普通情况:修改自己的预设 - preset_name = args_text - - if isinstance(event, GroupMessageEvent): - context_id = event.group_id - state = group_states[context_id] - is_group_target = True - else: # PrivateMessageEvent - if not plugin_config.enable_private_chat: - return - context_id = event.user_id - state = private_chat_states[context_id] - is_group_target = False - - # 处理关闭功能 - if preset_name == "off": - state.preset_name = preset_name - if target_id: - context_type = "群聊" if is_group_target else "私聊" - await preset_handler.finish(f"已关闭 {context_type} {context_id} 的llmchat功能") - elif isinstance(event, GroupMessageEvent): - await preset_handler.finish("已关闭llmchat群聊功能") - else: - await preset_handler.finish("已关闭llmchat私聊功能") - - # 检查预设是否存在 - if preset_name not in available_presets: - available_presets_str = "\n- ".join(available_presets) - await preset_handler.finish( - f"当前API预设:{state.preset_name}\n可用API预设:\n- {available_presets_str}" - ) - - # 切换预设 - state.preset_name = preset_name - if target_id: - context_type = "群聊" if is_group_target else "私聊" - await preset_handler.finish(f"已将 {context_type} {context_id} 切换至API预设:{preset_name}") - else: - await preset_handler.finish(f"已切换至API预设:{preset_name}") - - -edit_preset_handler = on_command( - "修改设定", - priority=1, - block=True, - permission=(SUPERUSER | GROUP_ADMIN | GROUP_OWNER | PRIVATE), +register_commands(plugin_config, states) +state_persistence = StatePersistence( + plugin_config, + states, + store.get_plugin_data_file("llmchat_state.json"), + store.get_plugin_data_file("llmchat_private_state.json"), ) -@edit_preset_handler.handle() -async def handle_edit_preset(event: GroupMessageEvent | PrivateMessageEvent, args: Message = CommandArg()): - if isinstance(event, GroupMessageEvent): - context_id = event.group_id - state = group_states[context_id] - else: # PrivateMessageEvent - if not plugin_config.enable_private_chat: - return - context_id = event.user_id - state = private_chat_states[context_id] - - group_prompt = args.extract_plain_text().strip() - state.group_prompt = group_prompt - await edit_preset_handler.finish("修改成功") - - -reset_handler = on_command( - "记忆清除", - priority=1, - block=True, - permission=(SUPERUSER | GROUP_ADMIN | GROUP_OWNER | PRIVATE), -) - - -@reset_handler.handle() -async def handle_reset(event: GroupMessageEvent | PrivateMessageEvent, args: Message = CommandArg()): - if isinstance(event, GroupMessageEvent): - context_id = event.group_id - state = group_states[context_id] - else: # PrivateMessageEvent - if not plugin_config.enable_private_chat: - return - context_id = event.user_id - state = private_chat_states[context_id] - - state.past_events.clear() - state.history.clear() - await reset_handler.finish("记忆已清空") - - -set_prob_handler = on_command( - "设置主动回复概率", - priority=1, - block=True, - permission=(SUPERUSER | GROUP_ADMIN | GROUP_OWNER), -) - - -@set_prob_handler.handle() -async def handle_set_prob(event: GroupMessageEvent, args: Message = CommandArg()): - context_id = event.group_id - state = group_states[context_id] - - try: - prob = float(args.extract_plain_text().strip()) - if prob < 0 or prob > 1: - raise ValueError("概率值必须在0-1之间") - except ValueError as e: - await set_prob_handler.finish(f"输入有误,请使用 [0,1] 的浮点数\n{e!s}") - return - - state.random_trigger_prob = prob - await set_prob_handler.finish(f"主动回复概率已设为 {prob}") - - -# 思维输出切换命令 -think_handler = on_command( - "切换思维输出", - priority=1, - block=True, - permission=(SUPERUSER | GROUP_ADMIN | GROUP_OWNER | PRIVATE), -) - - -@think_handler.handle() -async def handle_think(event: GroupMessageEvent | PrivateMessageEvent, args: Message = CommandArg()): - if isinstance(event, GroupMessageEvent): - state = group_states[event.group_id] - else: # PrivateMessageEvent - if not plugin_config.enable_private_chat: - return - state = private_chat_states[event.user_id] - - state.output_reasoning_content = not state.output_reasoning_content - - await think_handler.finish( - f"已{(state.output_reasoning_content and '开启') or '关闭'}思维输出" - ) - - -# region 持久化与定时任务 - -# 获取插件数据目录 -data_dir = store.get_plugin_data_dir() -# 获取插件数据文件 -data_file = store.get_plugin_data_file("llmchat_state.json") -private_data_file = store.get_plugin_data_file("llmchat_private_state.json") - - -async def save_state(): - """保存群组状态到文件""" - logger.info(f"开始保存群组状态到文件:{data_file}") - data = { - gid: { - "preset": state.preset_name, - "history": list(state.history), - "last_active": state.last_active, - "group_prompt": state.group_prompt, - "output_reasoning_content": state.output_reasoning_content, - "random_trigger_prob": state.random_trigger_prob, - } - for gid, state in group_states.items() - } - - os.makedirs(os.path.dirname(data_file), exist_ok=True) - async with aiofiles.open(data_file, "w", encoding="utf8") as f: - await f.write(json.dumps(data, ensure_ascii=False)) - - # 保存私聊状态 - if plugin_config.enable_private_chat: - logger.info(f"开始保存私聊状态到文件:{private_data_file}") - private_data = { - uid: { - "preset": state.preset_name, - "history": list(state.history), - "last_active": state.last_active, - "group_prompt": state.group_prompt, - "output_reasoning_content": state.output_reasoning_content, - } - for uid, state in private_chat_states.items() - } - - os.makedirs(os.path.dirname(private_data_file), exist_ok=True) - async with aiofiles.open(private_data_file, "w", encoding="utf8") as f: - await f.write(json.dumps(private_data, ensure_ascii=False)) - - -async def load_state(): - """从文件加载群组状态""" - logger.info(f"从文件加载群组状态:{data_file}") - if not os.path.exists(data_file): - return - - async with aiofiles.open(data_file, encoding="utf8") as f: - data = json.loads(await f.read()) - for gid, state_data in data.items(): - state = GroupState() - state.preset_name = state_data["preset"] - state.history = deque( - state_data["history"], maxlen=plugin_config.history_size * 2 - ) - state.last_active = state_data["last_active"] - state.group_prompt = state_data["group_prompt"] - state.output_reasoning_content = state_data["output_reasoning_content"] - state.random_trigger_prob = state_data.get("random_trigger_prob", plugin_config.random_trigger_prob) - group_states[int(gid)] = state - - # 加载私聊状态 - if plugin_config.enable_private_chat: - logger.info(f"从文件加载私聊状态:{private_data_file}") - if os.path.exists(private_data_file): - async with aiofiles.open(private_data_file, encoding="utf8") as f: - private_data = json.loads(await f.read()) - for uid, state_data in private_data.items(): - state = PrivateChatState() - state.preset_name = state_data["preset"] - state.history = deque( - state_data["history"], maxlen=plugin_config.history_size * 2 - ) - state.last_active = state_data["last_active"] - state.group_prompt = state_data["group_prompt"] - state.output_reasoning_content = state_data["output_reasoning_content"] - private_chat_states[int(uid)] = state - - -# 注册生命周期事件 @driver.on_startup -async def init_plugin(): - logger.info("插件启动初始化") - await load_state() - # 每5分钟保存状态 - scheduler.add_job(save_state, "interval", minutes=5) +async def init_plugin() -> None: + logger.info("llmchat插件启动初始化") + await state_persistence.load() + scheduler.add_job( + state_persistence.save, + "interval", + minutes=5, + id="llmchat_save_state", + replace_existing=True, + ) @driver.on_shutdown -async def cleanup_plugin(): - logger.info("插件关闭清理") - await save_state() - # 销毁MCPClient单例 +async def cleanup_plugin() -> None: + logger.info("llmchat插件关闭清理") + await message_dispatcher.shutdown() + await state_persistence.save() await MCPClient.destroy_instance() diff --git a/nonebot_plugin_llmchat/commands.py b/nonebot_plugin_llmchat/commands.py new file mode 100644 index 0000000..fc5ee63 --- /dev/null +++ b/nonebot_plugin_llmchat/commands.py @@ -0,0 +1,112 @@ +from nonebot import on_command +from nonebot.adapters.onebot.v11 import GroupMessageEvent, Message, PrivateMessageEvent +from nonebot.adapters.onebot.v11.permission import GROUP_ADMIN, GROUP_OWNER, PRIVATE +from nonebot.params import CommandArg +from nonebot.permission import SUPERUSER + +from .config import ScopedConfig +from .state import ChatState, StateStore + + +def _event_state(config: ScopedConfig, states: StateStore, event) -> ChatState | None: + if isinstance(event, GroupMessageEvent): + return states.group_states[event.group_id] + if config.enable_private_chat: + return states.private_states[event.user_id] + return None + + +def register_commands(config: ScopedConfig, states: StateStore) -> None: + preset_handler = on_command("API预设", priority=1, block=True, permission=SUPERUSER) + + @preset_handler.handle() + async def handle_preset(event: GroupMessageEvent | PrivateMessageEvent, args: Message = CommandArg()): + text = args.extract_plain_text().strip() + parts = text.split(maxsplit=1) + target_id: int | None = None + is_group_target = isinstance(event, GroupMessageEvent) + state = _event_state(config, states, event) + if state is None: + return + preset_name = text + if isinstance(event, PrivateMessageEvent) and parts and parts[0].isdigit(): + target_id = int(parts[0]) + is_group_target = target_id in states.group_states + state = states.get(target_id, is_group_target) + preset_name = parts[1] if len(parts) > 1 else "" + + available = {preset.name for preset in config.api_presets} + if not preset_name or preset_name not in available | {"off"}: + names = "\n- ".join(sorted(available)) + await preset_handler.finish(f"当前API预设:{state.preset_name}\n可用API预设:\n- {names}") + state.preset_name = preset_name + if target_id is not None: + kind = "群聊" if is_group_target else "私聊" + await preset_handler.finish(f"已将 {kind} {target_id} 切换至API预设:{preset_name}") + await preset_handler.finish(f"已切换至API预设:{preset_name}") + + edit_handler = on_command( + "修改设定", + priority=1, + block=True, + permission=SUPERUSER | GROUP_ADMIN | GROUP_OWNER | PRIVATE, + ) + + @edit_handler.handle() + async def handle_edit(event: GroupMessageEvent | PrivateMessageEvent, args: Message = CommandArg()): + state = _event_state(config, states, event) + if state is None: + return + state.prompt = args.extract_plain_text().strip() + await edit_handler.finish("修改成功") + + reset_handler = on_command( + "记忆清除", + priority=1, + block=True, + permission=SUPERUSER | GROUP_ADMIN | GROUP_OWNER | PRIVATE, + ) + + @reset_handler.handle() + async def handle_reset(event: GroupMessageEvent | PrivateMessageEvent): + state = _event_state(config, states, event) + if state is None: + return + state.pending_events.clear() + state.history.clear() + await reset_handler.finish("记忆已清空") + + probability_handler = on_command( + "设置主动回复概率", + priority=1, + block=True, + permission=SUPERUSER | GROUP_ADMIN | GROUP_OWNER, + ) + + @probability_handler.handle() + async def handle_probability(event: GroupMessageEvent, args: Message = CommandArg()): + try: + probability = float(args.extract_plain_text().strip()) + if not 0 <= probability <= 1: + raise ValueError("概率必须在0到1之间") + except ValueError as error: + await probability_handler.finish(f"输入有误,请使用 [0,1] 的浮点数\n{error!s}") + return + states.group_states[event.group_id].random_trigger_prob = probability + await probability_handler.finish(f"主动回复概率已设为 {probability}") + + think_handler = on_command( + "切换思维输出", + priority=1, + block=True, + permission=SUPERUSER | GROUP_ADMIN | GROUP_OWNER | PRIVATE, + ) + + @think_handler.handle() + async def handle_think(event: GroupMessageEvent | PrivateMessageEvent): + state = _event_state(config, states, event) + if state is None: + return + state.output_reasoning_content = not state.output_reasoning_content + status = "开启" if state.output_reasoning_content else "关闭" + await think_handler.finish(f"已{status}思维输出") diff --git a/nonebot_plugin_llmchat/config.py b/nonebot_plugin_llmchat/config.py index 834e109..4d957af 100755 --- a/nonebot_plugin_llmchat/config.py +++ b/nonebot_plugin_llmchat/config.py @@ -1,66 +1,82 @@ -from pydantic import BaseModel, Field +from pydantic import BaseModel, Field, model_validator class PresetConfig(BaseModel): - """API预设配置""" + """API预设配置。""" name: str = Field(..., description="预设名称(唯一标识)") api_base: str = Field(..., description="API基础地址") api_key: str = Field(..., description="API密钥") model_name: str = Field(..., description="模型名称") - max_tokens: int = Field(2048, description="最大响应token数") - temperature: float = Field(0.7, description="生成温度(0-2]") - proxy: str = Field("", description="HTTP代理服务器") - support_mcp: bool = Field(False, description="是否支持MCP") - support_image: bool = Field(False, description="是否支持图片输入") - extra_body: dict = Field({}, description="额外的请求体字段,用于兼容不同API的特殊参数") + max_tokens: int = Field(default=2048, description="最大响应token数") + temperature: float = Field(default=0.7, description="生成温度(0-2]") + proxy: str = Field(default="", description="HTTP代理服务器") + support_mcp: bool = Field(default=False, description="是否支持MCP") + support_image: bool = Field(default=False, description="是否支持图片输入") + extra_body: dict = Field(default_factory=dict, description="额外请求体字段") request_with_reasoning_content: bool = Field( - False, - description="请求中是否包含推理过程内容(部分模型要求进行了工具调用后,必须完整回传推理过程给API)" + default=False, + description="工具调用后是否向API回传推理内容", ) + class MCPServerConfig(BaseModel): - """MCP服务器配置""" - command: str | None = Field(None, description="stdio模式下MCP命令") - args: list[str] | None = Field([], description="stdio模式下MCP命令参数") - env: dict[str, str] | None = Field({}, description="stdio模式下MCP命令环境变量") - url: str | None = Field(None, description="远程MCP服务器地址") - headers: dict[str, str] | None = Field({}, description="远程MCP服务器http请求头,用于认证或其他设置") - transport: str | None = Field(None, description="远程MCP传输协议类型,可选 'sse' 或 'streamable_http',默认自动检测") + """MCP服务器配置。""" + + command: str | None = Field(default=None, description="stdio模式下MCP命令") + args: list[str] | None = Field(default_factory=list, description="stdio命令参数") + env: dict[str, str] | None = Field(default_factory=dict, description="stdio环境变量") + url: str | None = Field(default=None, description="远程MCP服务器地址") + headers: dict[str, str] | None = Field(default_factory=dict, description="HTTP请求头") + transport: str | None = Field(default=None, description="sse或streamable_http") + friendly_name: str | None = Field(default=None, description="MCP服务器友好名称") + additional_prompt: str | None = Field(default=None, description="额外提示词") + + @model_validator(mode="after") + def validate_transport(self): + if bool(self.command) == bool(self.url): + raise ValueError("MCP服务器必须且只能配置 command 或 url 其中之一") + if self.transport not in {None, "sse", "streamable_http"}: + raise ValueError("transport 必须是 sse 或 streamable_http") + return self - # 额外字段 - friendly_name: str | None = Field(None, description="MCP服务器友好名称") - additional_prompt: str | None = Field(None, description="额外提示词") class ScopedConfig(BaseModel): - """LLM Chat Plugin配置""" + """LLM Chat Plugin配置。""" - api_presets: list[PresetConfig] = Field( - ..., description="API预设列表(至少配置1个预设)" - ) - history_size: int = Field(20, description="LLM上下文消息保留数量") - past_events_size: int = Field(10, description="触发回复时发送的群消息数量") - request_timeout: int = Field(30, description="API请求超时时间(秒)") - default_preset: str = Field("off", description="默认使用的预设名称") - random_trigger_prob: float = Field( - 0.05, ge=0.0, le=1.0, description="随机触发概率(0-1]" - ) + api_presets: list[PresetConfig] = Field(..., description="API预设列表") + history_size: int = Field(default=20, ge=1, description="LLM上下文消息保留数量") + past_events_size: int = Field(default=10, ge=1, description="触发时发送的消息数量") + request_timeout: int = Field(default=30, ge=1, description="API请求超时时间(秒)") + max_tool_rounds: int = Field(default=8, ge=1, le=50, description="最大工具调用轮数") + max_repeated_tool_calls: int = Field(default=2, ge=1, le=10, description="相同工具调用最多执行次数") + mcp_timeout: int = Field(default=30, ge=1, description="MCP操作超时时间(秒)") + default_preset: str = Field(default="off", description="默认预设名称") + random_trigger_prob: float = Field(default=0.05, ge=0.0, le=1.0, description="随机触发概率") default_prompt: str = Field( - "你的回答应该尽量简洁、幽默、可以使用一些语气词、颜文字。你应该拒绝回答任何政治相关的问题。", + default="你的回答应该尽量简洁、幽默、可以使用一些语气词、颜文字。你应该拒绝回答任何政治相关的问题。", description="默认提示词", ) - mcp_server_cwd: str | None = Field( - None, - description="command类型MCP服务器的全局工作目录(cwd)" - ) - mcp_servers: dict[str, MCPServerConfig] = Field({}, description="MCP服务器配置") - blacklist_user_ids: set[int] = Field(set(), description="黑名单用户ID列表") - ignore_prefixes: list[str] = Field( - default_factory=list, - description="需要忽略的消息前缀列表,匹配到这些前缀的消息不会处理" - ) - enable_private_chat: bool = Field(False, description="是否启用私聊功能") - private_chat_preset: str = Field("off", description="私聊默认使用的预设名称") + mcp_server_cwd: str | None = Field(default=None, description="stdio MCP服务器全局工作目录") + mcp_servers: dict[str, MCPServerConfig] = Field(default_factory=dict, description="MCP服务器配置") + blacklist_user_ids: set[int] = Field(default_factory=set, description="黑名单用户ID") + ignore_prefixes: list[str] = Field(default_factory=list, description="忽略的消息前缀") + enable_private_chat: bool = Field(default=False, description="是否启用私聊") + private_chat_preset: str = Field(default="off", description="私聊默认预设") + + @model_validator(mode="after") + def validate_presets(self): + names = [preset.name for preset in self.api_presets] + if not names: + raise ValueError("api_presets 至少需要一个预设") + if len(names) != len(set(names)): + raise ValueError("api_presets 中的预设名称不能重复") + available = set(names) | {"off"} + if self.default_preset not in available: + raise ValueError(f"default_preset 不存在: {self.default_preset}") + if self.private_chat_preset not in available: + raise ValueError(f"private_chat_preset 不存在: {self.private_chat_preset}") + return self class Config(BaseModel): diff --git a/nonebot_plugin_llmchat/conversation.py b/nonebot_plugin_llmchat/conversation.py new file mode 100644 index 0000000..eec149c --- /dev/null +++ b/nonebot_plugin_llmchat/conversation.py @@ -0,0 +1,270 @@ +import asyncio +import base64 +from collections import Counter +import json +from typing import Any, cast + +import httpx +from nonebot import get_bot, logger +from nonebot.adapters.onebot.v11 import GroupMessageEvent, Message, MessageSegment +from openai import AsyncOpenAI + +from .config import PresetConfig, ScopedConfig +from .mcpclient import MCPClient +from .message_utils import ( + ChatEvent, + MessageSender, + build_reasoning_forward_nodes, + download_images, + format_message, + pop_reasoning_content, + send_split_messages, +) +from .prompts import build_system_prompt +from .state import ChatState, StateStore + + +class ConversationService: + def __init__( + self, + config: ScopedConfig, + states: StateStore, + bot_names: set[str], + sender: MessageSender, + ) -> None: + self.config = config + self.states = states + self.bot_names = bot_names + self.sender = sender + + async def process_event( + self, + context_id: int, + is_group: bool, + state: ChatState, + event: ChatEvent, + ) -> None: + snapshot = list(state.pending_events) + if not snapshot: + logger.debug(f"会话 {context_id} 没有待处理上下文,跳过重复触发") + return + state.pending_events.clear() + state.last_active = event.time + try: + await self._process_snapshot(context_id, is_group, state, event, snapshot) + except asyncio.CancelledError: + state.pending_events.extendleft(reversed(snapshot)) + raise + except Exception as error: + state.pending_events.extendleft(reversed(snapshot)) + logger.opt(exception=error).error(f"API请求失败 会话:{context_id}") + await self.sender(Message(f"服务暂时不可用,请稍后再试\n{error!s}")) + + def _create_client(self, preset: PresetConfig) -> AsyncOpenAI: + options: dict[str, Any] = { + "base_url": preset.api_base, + "api_key": preset.api_key, + "timeout": self.config.request_timeout, + } + if preset.proxy: + options["http_client"] = httpx.AsyncClient(proxy=preset.proxy) + return AsyncOpenAI(**options) + + async def _process_snapshot( + self, + context_id: int, + is_group: bool, + state: ChatState, + event: ChatEvent, + events: list[ChatEvent], + ) -> None: + preset = self.states.get_preset(state) + mcp_client = MCPClient.get_instance( + self.config.mcp_servers, + self.config.mcp_server_cwd, + self.config.mcp_timeout, + ) + system_prompt = build_system_prompt( + config=self.config, + state=state, + bot_names=self.bot_names, + is_group=is_group, + support_tools=preset.support_mcp, + ) + while state.history and state.history[0].get("role") != "user": + state.history.popleft() + messages: list[dict[str, Any]] = [ + {"role": "system", "content": system_prompt}, + *list(state.history)[-self.config.history_size * 2 :], + ] + content: list[dict[str, Any]] = [] + bot_name = next(iter(sorted(self.bot_names)), "机器人") + for pending_event in events: + content.append({"type": "text", "text": format_message(pending_event, bot_name)}) + if preset.support_image: + for image in await download_images(pending_event): + content.append({"type": "image_url", "image_url": {"url": f"data:image/jpeg;base64,{image}"}}) + transcript: list[dict[str, Any]] = [{"role": "user", "content": content}] + request_config: dict[str, Any] = { + "model": preset.model_name, + "max_tokens": preset.max_tokens, + "temperature": preset.temperature, + "timeout": self.config.request_timeout, + "extra_body": preset.extra_body, + } + if preset.support_mcp: + request_config["tools"] = await mcp_client.get_available_tools(is_group) + + client = self._create_client(preset) + try: + message = await self._run_tool_loop( + client=client, + preset=preset, + mcp_client=mcp_client, + request_config=request_config, + messages=messages, + transcript=transcript, + event=event, + is_group=is_group, + ) + finally: + await client.close() + + reply, tagged_reasoning = pop_reasoning_content(getattr(message, "content", None)) + reasoning = getattr(message, "reasoning_content", None) or tagged_reasoning + assistant_message: dict[str, Any] = {"role": "assistant", "content": reply or ""} + reply_images = getattr(message, "images", None) + if reply_images: + assistant_message["images"] = reply_images + if preset.request_with_reasoning_content: + assistant_message["reasoning_content"] = reasoning + transcript.append(assistant_message) + state.history.extend(transcript) + + if state.output_reasoning_content and reasoning: + await self._send_reasoning(context_id, is_group, event, reasoning) + if reply: + await send_split_messages(self.sender, reply) + await self._send_images(reply_images) + + async def _complete( + self, + client: AsyncOpenAI, + request_config: dict[str, Any], + messages: list[dict[str, Any]], + ) -> Any: + response = await cast(Any, client.chat.completions.create)(**request_config, messages=messages) + if not response.choices: + raise RuntimeError("API响应中没有choices") + if response.usage is not None: + logger.debug(f"API响应token数:{response.usage.total_tokens}") + return response.choices[0].message + + async def _run_tool_loop( + self, + *, + client: AsyncOpenAI, + preset: PresetConfig, + mcp_client: MCPClient, + request_config: dict[str, Any], + messages: list[dict[str, Any]], + transcript: list[dict[str, Any]], + event: ChatEvent, + is_group: bool, + ) -> Any: + message = await self._complete(client, request_config, messages + transcript) + if not preset.support_mcp: + return message + + call_counts: Counter[str] = Counter() + for round_number in range(1, self.config.max_tool_rounds + 1): + tool_calls = getattr(message, "tool_calls", None) + if not tool_calls: + return message + logger.info(f"处理第 {round_number}/{self.config.max_tool_rounds} 轮工具调用") + assistant_reply: dict[str, Any] = { + "role": "assistant", + "content": getattr(message, "content", None), + "tool_calls": [tool_call.model_dump() for tool_call in tool_calls], + } + reasoning = getattr(message, "reasoning_content", None) + if preset.request_with_reasoning_content and reasoning is not None: + assistant_reply["reasoning_content"] = reasoning + transcript.append(assistant_reply) + if getattr(message, "content", None): + await send_split_messages(self.sender, message.content) + + for tool_call in tool_calls: + await self._handle_tool_call( + mcp_client, + transcript, + tool_call, + call_counts, + event, + is_group, + ) + message = await self._complete(client, request_config, messages + transcript) + + if not getattr(message, "tool_calls", None): + return message + logger.warning(f"工具调用达到上限 {self.config.max_tool_rounds},强制要求模型总结") + final_config = {key: value for key, value in request_config.items() if key not in {"tools", "tool_choice"}} + final_messages = [ + *messages, + *transcript, + { + "role": "system", + "content": "工具调用次数已达上限。不得再调用工具,请根据已有结果直接给出最终回答。", + }, + ] + return await self._complete(client, final_config, final_messages) + + async def _handle_tool_call( + self, + mcp_client: MCPClient, + transcript: list[dict[str, Any]], + tool_call: Any, + call_counts: Counter[str], + event: ChatEvent, + is_group: bool, + ) -> None: + name = tool_call.function.name + raw_arguments = tool_call.function.arguments + try: + arguments = json.loads(raw_arguments) + if not isinstance(arguments, dict): + raise TypeError("arguments必须是JSON对象") + except (json.JSONDecodeError, TypeError, ValueError) as error: + result = f"工具参数格式错误: {error!s}" + else: + signature = f"{name}:{json.dumps(arguments, sort_keys=True, ensure_ascii=False)}" + call_counts[signature] += 1 + if call_counts[signature] > self.config.max_repeated_tool_calls: + result = "相同工具和参数的调用已达到上限,请使用已有结果继续回答。" + logger.warning(f"阻止重复工具调用: {signature}") + else: + await self.sender(Message(f"正在使用{mcp_client.get_friendly_name(name)}")) + result = await mcp_client.call_tool( + name, + arguments, + group_id=event.group_id if is_group and isinstance(event, GroupMessageEvent) else None, + bot_id=str(event.self_id), + ) + transcript.append({"role": "tool", "tool_call_id": tool_call.id, "content": str(result)}) + + async def _send_reasoning(self, context_id: int, is_group: bool, event: ChatEvent, content: str) -> None: + try: + bot = get_bot(str(event.self_id)) + nickname = next(iter(sorted(self.bot_names)), "机器人") + nodes = build_reasoning_forward_nodes(bot.self_id, nickname, content) + if is_group: + await bot.send_group_forward_msg(group_id=context_id, messages=nodes) + else: + await bot.send_private_forward_msg(user_id=context_id, messages=nodes) + except Exception: + logger.exception("合并转发思维内容失败") + + async def _send_images(self, images: Any) -> None: + for image in images or []: + encoded = image["image_url"]["url"].split(",", maxsplit=1)[-1] + await self.sender(Message(MessageSegment.image(base64.b64decode(encoded)))) diff --git a/nonebot_plugin_llmchat/dispatcher.py b/nonebot_plugin_llmchat/dispatcher.py new file mode 100644 index 0000000..bd06923 --- /dev/null +++ b/nonebot_plugin_llmchat/dispatcher.py @@ -0,0 +1,66 @@ +import asyncio +from collections.abc import Awaitable, Callable +from typing import Any + +from nonebot import logger + +from .state import ChatState + +EventProcessor = Callable[[int, bool, ChatState, Any], Awaitable[None]] + + +class MessageDispatcher: + """每个会话只运行一个worker,并可靠接管竞态窗口中新入队的消息。""" + + def __init__(self, processor: EventProcessor): + self._processor = processor + self._tasks: set[asyncio.Task[None]] = set() + self._closing = False + + async def enqueue(self, context_id: int, is_group: bool, state: ChatState, event: Any) -> None: + if self._closing: + return + await state.queue.put(event) + async with state.worker_lock: + if state.worker_task is None or state.worker_task.done(): + self._start_worker(context_id, is_group, state) + + def _start_worker(self, context_id: int, is_group: bool, state: ChatState) -> None: + task = asyncio.create_task( + self._run_worker(context_id, is_group, state), + name=f"llmchat:{'group' if is_group else 'private'}:{context_id}", + ) + state.worker_task = task + self._tasks.add(task) + task.add_done_callback(self._tasks.discard) + + async def _run_worker(self, context_id: int, is_group: bool, state: ChatState) -> None: + current_task = asyncio.current_task() + try: + while True: + try: + event = state.queue.get_nowait() + except asyncio.QueueEmpty: + break + try: + await self._processor(context_id, is_group, state, event) + except asyncio.CancelledError: + raise + except Exception: + logger.exception(f"处理会话 {context_id} 的消息失败") + finally: + state.queue.task_done() + finally: + async with state.worker_lock: + if state.worker_task is current_task: + state.worker_task = None + if not self._closing and not state.queue.empty(): + self._start_worker(context_id, is_group, state) + + async def shutdown(self) -> None: + self._closing = True + tasks = list(self._tasks) + for task in tasks: + task.cancel() + if tasks: + await asyncio.gather(*tasks, return_exceptions=True) diff --git a/nonebot_plugin_llmchat/mcpclient.py b/nonebot_plugin_llmchat/mcpclient.py index ada9c6e..ea7e10f 100644 --- a/nonebot_plugin_llmchat/mcpclient.py +++ b/nonebot_plugin_llmchat/mcpclient.py @@ -24,6 +24,7 @@ class MCPClient: cls, server_config: dict[str, MCPServerConfig] | None = None, default_command_cwd: str | None = None, + operation_timeout: int = 30, ): if cls._instance is None: cls._instance = super().__new__(cls) @@ -33,6 +34,7 @@ class MCPClient: self, server_config: dict[str, MCPServerConfig] | None = None, default_command_cwd: str | None = None, + operation_timeout: int = 30, ): if self._initialized: return @@ -43,6 +45,7 @@ class MCPClient: logger.info(f"正在初始化MCPClient单例,共有{len(server_config)}个服务器配置") self.server_config = server_config self.default_command_cwd = default_command_cwd + self.operation_timeout = operation_timeout self.sessions = {} self.exit_stack = AsyncExitStack() self._session_exit_stacks: dict[str, AsyncExitStack] = {} @@ -62,12 +65,13 @@ class MCPClient: cls, server_config: dict[str, MCPServerConfig] | None = None, default_command_cwd: str | None = None, + operation_timeout: int = 30, ): """获取MCPClient实例""" if cls._instance is None: if server_config is None: raise ValueError("server_config must be provided for first initialization") - cls._instance = cls(server_config, default_command_cwd) + cls._instance = cls(server_config, default_command_cwd, operation_timeout) return cls._instance @classmethod @@ -85,10 +89,24 @@ class MCPClient: await self._get_or_create_session(server_name) logger.info(f"已成功连接到MCP服务器[{server_name}]") - async def _create_server_session(self, server_name: str) -> tuple[ClientSession, AsyncExitStack]: + async def _open_server_session(self, server_name: str) -> tuple[ClientSession, AsyncExitStack]: + session_stack = AsyncExitStack() + try: + return await self._initialize_server_session(server_name, session_stack) + except BaseException: + try: + await asyncio.wait_for(session_stack.aclose(), timeout=self.operation_timeout) + except (asyncio.TimeoutError, RuntimeError): + logger.error(f"清理未完成的MCP会话[{server_name}]失败或超时") + raise + + async def _initialize_server_session( + self, + server_name: str, + session_stack: AsyncExitStack, + ) -> tuple[ClientSession, AsyncExitStack]: """创建并初始化一个新的服务器会话。""" config = self.server_config[server_name] - session_stack = AsyncExitStack() if config.url: transport_type = config.transport if transport_type == "streamable_http": @@ -137,6 +155,15 @@ class MCPClient: await session.initialize() return session, session_stack + async def _create_server_session(self, server_name: str) -> tuple[ClientSession, AsyncExitStack]: + try: + return await asyncio.wait_for( + self._open_server_session(server_name), + timeout=self.operation_timeout, + ) + except asyncio.TimeoutError as error: + raise TimeoutError(f"连接MCP服务器[{server_name}]超时") from error + async def _close_server_session(self, server_name: str): """关闭指定服务器会话。""" session_stack = self._session_exit_stacks.pop(server_name, None) @@ -144,7 +171,10 @@ class MCPClient: self._session_last_used.pop(server_name, None) if session_stack is not None: - await session_stack.aclose() + try: + await asyncio.wait_for(session_stack.aclose(), timeout=self.operation_timeout) + except asyncio.TimeoutError: + logger.error(f"关闭MCP会话[{server_name}]超时") async def _get_or_create_session(self, server_name: str) -> ClientSession: """获取可复用会话;若不存在或已过期则新建。""" @@ -204,7 +234,7 @@ class MCPClient: for server_name in self.server_config.keys(): logger.debug(f"正在从服务器[{server_name}]获取工具列表") session = await self._get_or_create_session(server_name) - response = await session.list_tools() + response = await asyncio.wait_for(session.list_tools(), timeout=self.operation_timeout) tools = response.tools logger.debug(f"在服务器[{server_name}]中找到{len(tools)}个工具") @@ -243,7 +273,13 @@ class MCPClient: if group_id is None or bot_id is None: return "QQ工具需要提供group_id和bot_id参数" logger.info(f"调用OneBot工具[{tool_name}]") - return await self.onebot_tools.call_tool(tool_name, tool_args, group_id, bot_id) + try: + return await asyncio.wait_for( + self.onebot_tools.call_tool(tool_name, tool_args, group_id, bot_id), + timeout=self.operation_timeout, + ) + except asyncio.TimeoutError: + return f"调用OneBot工具[{tool_name}]超时" # 检查是否是MCP工具 if tool_name.startswith("mcp__"): @@ -254,12 +290,14 @@ class MCPClient: server_name = parts[1] real_tool_name = parts[2] + if server_name not in self.server_config: + return f"未知的MCP服务器: {server_name}" logger.info(f"按需连接到服务器[{server_name}]调用工具[{real_tool_name}]") try: await self._ensure_cleanup_task() session = await self._get_or_create_session(server_name) - response = await asyncio.wait_for(session.call_tool(real_tool_name, tool_args), timeout=30) + response = await asyncio.wait_for(session.call_tool(real_tool_name, tool_args), timeout=self.operation_timeout) logger.debug(f"工具[{real_tool_name}]调用完成,响应: {response}") return response.content except asyncio.TimeoutError: @@ -289,7 +327,8 @@ class MCPClient: server_name = parts[1] real_tool_name = parts[2] - return (self.server_config[server_name].friendly_name or server_name) + " - " + real_tool_name + server = self.server_config.get(server_name) + return ((server.friendly_name if server else None) or server_name) + " - " + real_tool_name # 未知工具类型,返回原名称 return tool_name diff --git a/nonebot_plugin_llmchat/message_utils.py b/nonebot_plugin_llmchat/message_utils.py new file mode 100644 index 0000000..5f7a346 --- /dev/null +++ b/nonebot_plugin_llmchat/message_utils.py @@ -0,0 +1,93 @@ +import asyncio +import base64 +from collections.abc import Awaitable, Callable +from datetime import datetime +import json +import re +import ssl + +import httpx +from nonebot import logger +from nonebot.adapters.onebot.v11 import GroupMessageEvent, Message, PrivateMessageEvent + +ChatEvent = GroupMessageEvent | PrivateMessageEvent +MessageSender = Callable[[Message], Awaitable[object]] + + +def pop_reasoning_content(content: str | None) -> tuple[str | None, str | None]: + if content is None: + return None, None + if matched := re.match(r"(.*?)", content, flags=re.DOTALL): + return content.replace(matched.group(0), "").strip(), matched.group(1).strip() + return content, None + + +def format_message(event: ChatEvent, bot_name: str) -> str: + text = "" + if isinstance(event, GroupMessageEvent) and event.reply is not None: + text += f"[回复 {event.reply.sender.nickname} 的消息 {event.reply.message.extract_plain_text()}]\n" + if isinstance(event, GroupMessageEvent) and event.is_tome(): + text += f"@{bot_name} " + for segment in event.get_message(): + if segment.type == "at": + text += segment.data.get("name", "") + elif segment.type == "image": + text += "[图片]" + elif segment.type == "voice": + text += "[语音]" + elif segment.type == "text": + text += segment.data.get("text", "") + nickname = event.sender.card or event.sender.nickname if isinstance(event, GroupMessageEvent) else event.sender.nickname + return json.dumps( + { + "SenderNickname": str(nickname), + "SenderUserId": str(event.user_id), + "Message": text, + "MessageID": event.message_id, + "SendTime": datetime.fromtimestamp(event.time).isoformat(), + }, + ensure_ascii=False, + ) + + +async def download_images(event: ChatEvent) -> list[str]: + urls = [segment.data.get("url") or segment.data.get("file") for segment in event.get_message() if segment.type == "image"] + urls = [url for url in urls if url] + if not urls: + return [] + ssl_context = ssl.create_default_context() + ssl_context.check_hostname = False + ssl_context.verify_mode = ssl.CERT_NONE + ssl_context.set_ciphers("DEFAULT@SECLEVEL=2") + images: list[str] = [] + async with httpx.AsyncClient(verify=ssl_context, timeout=10.0) as client: + for url in urls: + try: + response = await client.get(url) + response.raise_for_status() + images.append(base64.b64encode(response.content).decode()) + except Exception: + logger.exception(f"下载图片失败: {url}") + return images + + +async def send_split_messages(sender: MessageSender, content: str) -> None: + segments = [segment.strip() for segment in content.split("") if segment.strip()] + logger.info(f"准备发送分段消息,分段数:{len(segments)}") + for index, segment in enumerate(segments): + if index: + await asyncio.sleep(2) + await sender(Message(segment)) + + +def build_reasoning_forward_nodes(self_id: str, nickname: str, content: str) -> list[dict]: + return [ + { + "type": "node", + "data": {"nickname": nickname, "user_id": self_id, "content": f"{nickname}的内心OS:"}, + }, + { + "type": "node", + "data": {"nickname": nickname, "user_id": self_id, "content": content}, + }, + ] diff --git a/nonebot_plugin_llmchat/persistence.py b/nonebot_plugin_llmchat/persistence.py new file mode 100644 index 0000000..c7f7a34 --- /dev/null +++ b/nonebot_plugin_llmchat/persistence.py @@ -0,0 +1,88 @@ +import asyncio +from collections import deque +import json +import os +from pathlib import Path + +import aiofiles +from nonebot import logger + +from .config import ScopedConfig +from .state import ChatState, StateStore + + +class StatePersistence: + def __init__(self, config: ScopedConfig, states: StateStore, group_file: Path, private_file: Path): + self.config = config + self.states = states + self.group_file = group_file + self.private_file = private_file + self._lock = asyncio.Lock() + + @staticmethod + def _serialize(state: ChatState, include_probability: bool) -> dict: + data = { + "preset": state.preset_name, + "history": list(state.history), + "last_active": state.last_active, + "group_prompt": state.prompt, + "output_reasoning_content": state.output_reasoning_content, + } + if include_probability: + data["random_trigger_prob"] = state.random_trigger_prob + return data + + async def _write_json(self, path: Path, data: dict) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + temporary = path.with_suffix(path.suffix + ".tmp") + async with aiofiles.open(temporary, "w", encoding="utf8") as file: + await file.write(json.dumps(data, ensure_ascii=False)) + os.replace(temporary, path) + + async def save(self) -> None: + async with self._lock: + await self._write_json( + self.group_file, + {key: self._serialize(state, True) for key, state in self.states.group_states.items()}, + ) + if self.config.enable_private_chat: + await self._write_json( + self.private_file, + {key: self._serialize(state, False) for key, state in self.states.private_states.items()}, + ) + + async def _read_json(self, path: Path) -> dict: + if not path.exists(): + return {} + try: + async with aiofiles.open(path, encoding="utf8") as file: + value = json.loads(await file.read()) + return value if isinstance(value, dict) else {} + except (OSError, json.JSONDecodeError): + logger.exception(f"读取llmchat状态失败: {path}") + return {} + + def _restore(self, raw: dict, *, is_group: bool) -> ChatState: + state = ChatState( + preset_name=raw.get( + "preset", + self.config.default_preset if is_group else self.config.private_chat_preset, + ), + history_size=self.config.history_size, + past_events_size=self.config.past_events_size, + random_trigger_prob=raw.get("random_trigger_prob", self.config.random_trigger_prob) if is_group else 0.0, + ) + state.history = deque(raw.get("history", []), maxlen=self.config.history_size * 2) + state.last_active = raw.get("last_active", state.last_active) + state.prompt = raw.get("group_prompt", raw.get("prompt")) + state.output_reasoning_content = raw.get("output_reasoning_content", False) + return state + + async def load(self) -> None: + groups = await self._read_json(self.group_file) + for context_id, raw in groups.items(): + self.states.group_states[int(context_id)] = self._restore(raw, is_group=True) + if self.config.enable_private_chat: + private = await self._read_json(self.private_file) + for context_id, raw in private.items(): + self.states.private_states[int(context_id)] = self._restore(raw, is_group=False) diff --git a/nonebot_plugin_llmchat/prompts.py b/nonebot_plugin_llmchat/prompts.py new file mode 100644 index 0000000..ba7e235 --- /dev/null +++ b/nonebot_plugin_llmchat/prompts.py @@ -0,0 +1,43 @@ +from .config import ScopedConfig +from .state import ChatState + + +def build_system_prompt( + *, + config: ScopedConfig, + state: ChatState, + bot_names: set[str], + is_group: bool, + support_tools: bool, +) -> str: + chat_type = "群聊" if is_group else "私聊" + names = "、".join(sorted(bot_names)) + lines = [ + f"我想要你帮我在{chat_type}中闲聊,大家一般叫你{names}。", + "我会在后面的信息中告诉你每条信息的发送者和发送时间,你可以直接称呼发送者的昵称。", + "你的回复需要遵守以下规则:", + "- 多条消息之间使用分隔,标记前后不需要额外换行或空格。", + "- 除外,不要输出其他类似标记。", + "- 不要使用Markdown或HTML,换行请直接使用换行符。", + "- 以普通人的方式聊天,每条消息尽量简短;代码应放在单独一条消息中。", + "- 只在第一次回答某位发送者时礼貌问候。", + "- 使用[CQ:reply,id=消息id]引用消息。", + "- 优先回复提到你的新消息;也可以选择不回复。", + "- 完全不回复时只输出。", + "- 尽量减少不必要的思考。", + ] + if is_group: + lines.append("- 使用[CQ:at,qq=QQ号]提及群成员,是否提及由你决定。") + lines.extend( + [ + "下面是你的性格设定;若其中指定了身份或名字,应优先遵守:", + state.prompt or config.default_prompt, + ] + ) + if support_tools: + additions = [ + f"{name}:{server.additional_prompt}" for name, server in config.mcp_servers.items() if server.additional_prompt + ] + if additions: + lines.extend(["你可以使用工具,额外说明如下:", *additions]) + return "\n".join(lines) diff --git a/nonebot_plugin_llmchat/state.py b/nonebot_plugin_llmchat/state.py new file mode 100644 index 0000000..038317e --- /dev/null +++ b/nonebot_plugin_llmchat/state.py @@ -0,0 +1,58 @@ +import asyncio +from collections import defaultdict, deque +from dataclasses import dataclass, field +import time +from typing import Any + +from .config import PresetConfig, ScopedConfig + + +@dataclass +class ChatState: + preset_name: str + history_size: int + past_events_size: int + random_trigger_prob: float = 0.0 + history: deque[dict[str, Any]] = field(init=False) + pending_events: deque[Any] = field(init=False) + queue: asyncio.Queue[Any] = field(default_factory=asyncio.Queue) + worker_lock: asyncio.Lock = field(default_factory=asyncio.Lock) + worker_task: asyncio.Task[None] | None = None + last_active: float = field(default_factory=time.time) + prompt: str | None = None + output_reasoning_content: bool = False + + def __post_init__(self) -> None: + self.history = deque(maxlen=self.history_size * 2) + self.pending_events = deque(maxlen=self.past_events_size) + + +class StateStore: + def __init__(self, config: ScopedConfig): + self.config = config + self.group_states: dict[int, ChatState] = defaultdict(self._new_group_state) + self.private_states: dict[int, ChatState] = defaultdict(self._new_private_state) + + def _new_group_state(self) -> ChatState: + return ChatState( + self.config.default_preset, + self.config.history_size, + self.config.past_events_size, + self.config.random_trigger_prob, + ) + + def _new_private_state(self) -> ChatState: + return ChatState( + self.config.private_chat_preset, + self.config.history_size, + self.config.past_events_size, + ) + + def get(self, context_id: int, is_group: bool) -> ChatState: + return self.group_states[context_id] if is_group else self.private_states[context_id] + + def get_preset(self, state: ChatState) -> PresetConfig: + return next( + (preset for preset in self.config.api_presets if preset.name == state.preset_name), + self.config.api_presets[0], + ) diff --git a/tests/__init__.py b/tests/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/bootstrap.py b/tests/bootstrap.py new file mode 100644 index 0000000..b7b321b --- /dev/null +++ b/tests/bootstrap.py @@ -0,0 +1,11 @@ +"""加载子模块而不执行NoneBot插件注册入口。""" + +from pathlib import Path +import sys +from types import ModuleType + +PACKAGE_NAME = "nonebot_plugin_llmchat" +if PACKAGE_NAME not in sys.modules: + package = ModuleType(PACKAGE_NAME) + package.__path__ = [str(Path(__file__).parents[1] / PACKAGE_NAME)] + sys.modules[PACKAGE_NAME] = package diff --git a/tests/test_config_and_messages.py b/tests/test_config_and_messages.py new file mode 100644 index 0000000..2b9826d --- /dev/null +++ b/tests/test_config_and_messages.py @@ -0,0 +1,31 @@ +# ruff: noqa: I001 +import unittest + +from pydantic import ValidationError + +import tests.bootstrap # noqa: F401 +from nonebot_plugin_llmchat.config import MCPServerConfig, PresetConfig, ScopedConfig +from nonebot_plugin_llmchat.message_utils import pop_reasoning_content + + +def assert_validation_error(factory): + try: + factory() + except ValidationError: + return + raise AssertionError("expected ValidationError") + + +class ConfigurationTests(unittest.TestCase): + def test_rejects_duplicate_presets(self): + preset = PresetConfig(name="same", api_base="x", api_key="x", model_name="x") + assert_validation_error(lambda: ScopedConfig(api_presets=[preset, preset])) + + def test_mcp_requires_exactly_one_transport_target(self): + assert_validation_error(MCPServerConfig) + assert_validation_error(lambda: MCPServerConfig(command="cmd", url="https://example.invalid")) + + def test_reasoning_tag_is_removed(self): + reply, reasoning = pop_reasoning_content("secretanswer") + assert reply == "answer" + assert reasoning == "secret" diff --git a/tests/test_dispatcher.py b/tests/test_dispatcher.py new file mode 100644 index 0000000..c7ec180 --- /dev/null +++ b/tests/test_dispatcher.py @@ -0,0 +1,49 @@ +import asyncio +import unittest + +# ruff: noqa: I001 +import tests.bootstrap # noqa: F401 +from nonebot_plugin_llmchat.dispatcher import MessageDispatcher +from nonebot_plugin_llmchat.state import ChatState + + +class DispatcherTests(unittest.IsolatedAsyncioTestCase): + async def test_messages_enqueued_while_worker_runs_are_not_stranded(self): + started = asyncio.Event() + release = asyncio.Event() + processed = [] + + async def process(_context_id, _is_group, _state, event): + processed.append(event) + if event == "first": + started.set() + await release.wait() + + state = ChatState("test", 10, 10) + dispatcher = MessageDispatcher(process) + await dispatcher.enqueue(1, True, state, "first") + await asyncio.wait_for(started.wait(), timeout=1) + await dispatcher.enqueue(1, True, state, "second") + release.set() + await asyncio.wait_for(state.queue.join(), timeout=1) + + assert processed, ["first", "second"] + assert state.queue.empty() + await dispatcher.shutdown() + + async def test_processor_failure_does_not_stop_the_queue(self): + processed = [] + + async def process(_context_id, _is_group, _state, event): + processed.append(event) + if event == "bad": + raise RuntimeError("boom") + + state = ChatState("test", 10, 10) + dispatcher = MessageDispatcher(process) + await dispatcher.enqueue(1, True, state, "bad") + await dispatcher.enqueue(1, True, state, "good") + await asyncio.wait_for(state.queue.join(), timeout=1) + + assert processed, ["bad", "good"] + await dispatcher.shutdown() diff --git a/tests/test_persistence.py b/tests/test_persistence.py new file mode 100644 index 0000000..120381c --- /dev/null +++ b/tests/test_persistence.py @@ -0,0 +1,34 @@ +# ruff: noqa: I001 +import tempfile +from pathlib import Path +import unittest + +import tests.bootstrap # noqa: F401 +from nonebot_plugin_llmchat.config import PresetConfig, ScopedConfig +from nonebot_plugin_llmchat.persistence import StatePersistence +from nonebot_plugin_llmchat.state import StateStore + + +class PersistenceTests(unittest.IsolatedAsyncioTestCase): + async def test_private_state_loads_even_when_group_file_is_missing(self): + preset = PresetConfig(name="test", api_base="x", api_key="x", model_name="x") + config = ScopedConfig( + api_presets=[preset], + default_preset="test", + enable_private_chat=True, + private_chat_preset="test", + ) + with tempfile.TemporaryDirectory() as directory: + group_file = Path(directory) / "groups.json" + private_file = Path(directory) / "private.json" + source = StateStore(config) + source.private_states[42].prompt = "remember me" + persistence = StatePersistence(config, source, group_file, private_file) + await persistence.save() + group_file.unlink() + + restored = StateStore(config) + await StatePersistence(config, restored, group_file, private_file).load() + + assert restored.private_states[42].prompt == "remember me" + assert not group_file.exists() diff --git a/tests/test_tool_loop.py b/tests/test_tool_loop.py new file mode 100644 index 0000000..b21f948 --- /dev/null +++ b/tests/test_tool_loop.py @@ -0,0 +1,101 @@ +from types import SimpleNamespace +from typing import Any, cast +import unittest + +# ruff: noqa: I001 +import tests.bootstrap # noqa: F401 +from nonebot_plugin_llmchat.config import PresetConfig, ScopedConfig +from nonebot_plugin_llmchat.conversation import ConversationService +from nonebot_plugin_llmchat.state import StateStore + + +class FakeToolCall: + def __init__(self): + self.id = "call-1" + self.function = SimpleNamespace(name="mcp__demo__search", arguments='{"query":"same"}') + + def model_dump(self): + return { + "id": self.id, + "type": "function", + "function": {"name": self.function.name, "arguments": self.function.arguments}, + } + + +class FakeMessage: + def __init__(self, content=None, tool_calls=None): + self.content = content + self.tool_calls = tool_calls + self.reasoning_content = None + + +class FakeCompletions: + def __init__(self): + self.calls = [] + + async def create(self, **kwargs): + self.calls.append(kwargs) + if "tools" not in kwargs: + message = FakeMessage("final answer") + else: + message = FakeMessage(tool_calls=[FakeToolCall()]) + return SimpleNamespace(choices=[SimpleNamespace(message=message)], usage=None) + + +class FakeClient: + def __init__(self): + self.chat = SimpleNamespace(completions=FakeCompletions()) + + +class FakeMCPClient: + def __init__(self): + self.call_count = 0 + + def get_friendly_name(self, _name): + return "测试工具" + + async def call_tool(self, *_args, **_kwargs): + self.call_count += 1 + return "result" + + +class ToolLoopTests(unittest.IsolatedAsyncioTestCase): + async def test_repeated_tool_calls_are_bounded_and_forced_to_finish(self): + preset = PresetConfig( + name="test", + api_base="https://example.invalid/v1", + api_key="test", + model_name="test", + support_mcp=True, + ) + config = ScopedConfig( + api_presets=[preset], + default_preset="test", + max_tool_rounds=3, + max_repeated_tool_calls=2, + ) + sent = [] + + async def sender(message): + sent.append(message) + + service = ConversationService(config, StateStore(config), {"bot"}, sender) + client = FakeClient() + mcp_client = FakeMCPClient() + transcript = [{"role": "user", "content": "hello"}] + message = await service._run_tool_loop( + client=cast(Any, client), + preset=preset, + mcp_client=cast(Any, mcp_client), + request_config={"model": "test", "tools": [{}]}, + messages=[{"role": "system", "content": "system"}], + transcript=transcript, + event=cast(Any, SimpleNamespace(self_id=1)), + is_group=False, + ) + + assert message.content == "final answer" + assert mcp_client.call_count == 2 + assert len(client.chat.completions.calls) == 5 + assert "tools" not in client.chat.completions.calls[-1] + assert "达到上限" in transcript[-1]["content"]