# 微信公众号的加解密算法与企业微信一样,所以直接使用企业微信的加解密算法文件 import time import traceback from langbot.libs.wecom_api.WXBizMsgCrypt3 import WXBizMsgCrypt import xml.etree.ElementTree as ET from quart import Quart, request import hashlib from typing import Callable from langbot.libs.official_account_api.oaevent import OAEvent import asyncio xml_template = """ {create_time} """ _MAX_CALLBACK_BODY_BYTES = 1024 * 1024 class OAClient: _STATE_TTL_SECONDS = 600 _STATE_MAX = 4096 _MAX_CONTENT_CHARS = 200000 def __init__( self, token: str, EncodingAESKey: str, AppID: str, Appsecret: str, logger: None, unified_mode: bool = False, api_base_url: str = 'https://api.weixin.qq.com', ): self.token = token self.aes = EncodingAESKey self.appid = AppID self.appsecret = Appsecret self.base_url = api_base_url self.access_token = '' self.unified_mode = unified_mode self.app = Quart(__name__) self.app.config['MAX_CONTENT_LENGTH'] = _MAX_CALLBACK_BODY_BYTES # 只有在非统一模式下才注册独立路由 if not self.unified_mode: self.app.add_url_rule( '/callback/command', 'handle_callback', self.handle_callback_request, methods=['GET', 'POST'], ) self._message_handlers = { 'example': [], } self.access_token_expiry_time = None self.msg_id_map = {} self.generated_content = {} self._msg_seen_at = {} self._generated_at = {} self._last_state_prune = 0.0 self.logger = logger def _prune_state(self) -> None: now = time.monotonic() if now - self._last_state_prune >= 60: self._last_state_prune = now for message_id, seen_at in tuple(self._msg_seen_at.items()): if now - seen_at > self._STATE_TTL_SECONDS: self._msg_seen_at.pop(message_id, None) self.msg_id_map.pop(message_id, None) for message_id, generated_at in tuple(self._generated_at.items()): if now - generated_at > self._STATE_TTL_SECONDS: self._generated_at.pop(message_id, None) self.generated_content.pop(message_id, None) while len(self.msg_id_map) > self._STATE_MAX: message_id = next(iter(self.msg_id_map)) self.msg_id_map.pop(message_id, None) self._msg_seen_at.pop(message_id, None) while len(self.generated_content) > self._STATE_MAX: message_id = next(iter(self.generated_content)) self.generated_content.pop(message_id, None) self._generated_at.pop(message_id, None) def clear(self) -> None: self.msg_id_map.clear() self.generated_content.clear() self._msg_seen_at.clear() self._generated_at.clear() async def handle_callback_request(self): """处理回调请求(独立端口模式,使用全局 request)。""" return await self._handle_callback_internal(request) async def handle_unified_webhook(self, req): """处理回调请求(统一 webhook 模式,显式传递 request)。 Args: req: Quart Request 对象 Returns: 响应数据 """ return await self._handle_callback_internal(req) async def _handle_callback_internal(self, req): """处理回调请求的内部实现,包括 GET 验证和 POST 消息接收。 Args: req: Quart Request 对象 """ try: # 每隔100毫秒查询是否生成ai回答 start_time = time.time() signature = req.args.get('signature', '') timestamp = req.args.get('timestamp', '') nonce = req.args.get('nonce', '') echostr = req.args.get('echostr', '') msg_signature = req.args.get('msg_signature', '') if msg_signature is None: await self.logger.error('msg_signature不在请求体中') raise Exception('msg_signature不在请求体中') if req.method == 'GET': # 校验签名 check_str = ''.join(sorted([self.token, timestamp, nonce])) check_signature = hashlib.sha1(check_str.encode('utf-8')).hexdigest() if check_signature == signature: return echostr # 验证成功返回echostr else: await self.logger.error('拒绝请求') raise Exception('拒绝请求') elif req.method == 'POST': encryt_msg = await req.data if len(encryt_msg) > _MAX_CALLBACK_BODY_BYTES: raise ValueError('Official Account callback body exceeds the size limit') wxcpt = WXBizMsgCrypt(self.token, self.aes, self.appid) ret, xml_msg = await asyncio.to_thread( wxcpt.DecryptMsg, encryt_msg, msg_signature, timestamp, nonce, ) xml_msg = xml_msg.decode('utf-8') if ret != 0: await self.logger.error('消息解密失败') raise Exception('消息解密失败') message_data = await self.get_message(xml_msg) if message_data: event = OAEvent.from_payload(message_data) if event: await self._handle_message(event) root = await asyncio.to_thread(ET.fromstring, xml_msg) from_user = root.find('FromUserName').text # 发送者 to_user = root.find('ToUserName').text # 机器人 timeout = 4.80 interval = 0.1 while True: content = self.generated_content.pop(message_data['MsgId'], None) self._generated_at.pop(message_data['MsgId'], None) if content: response_xml = xml_template.format( to_user=from_user, from_user=to_user, create_time=int(time.time()), content=content, ) return response_xml if time.time() - start_time >= timeout: break await asyncio.sleep(interval) if self.msg_id_map.get(message_data['MsgId'], 1) == 3: # response_xml = xml_template.format( # to_user=from_user, # from_user=to_user, # create_time=int(time.time()), # content = "请求失效:暂不支持公众号超过15秒的请求,如有需求,请联系 LangBot 团队。" # ) print('请求失效:暂不支持公众号超过15秒的请求,如有需求,请联系 LangBot 团队。') return '' except Exception: await self.logger.error(f'handle_callback_request失败: {traceback.format_exc()}') traceback.print_exc() async def get_message(self, xml_msg: str): root = await asyncio.to_thread(ET.fromstring, xml_msg) message_data = { 'ToUserName': root.find('ToUserName').text, 'FromUserName': root.find('FromUserName').text, 'CreateTime': int(root.find('CreateTime').text), 'MsgType': root.find('MsgType').text, 'Content': root.find('Content').text if root.find('Content') is not None else None, 'MsgId': int(root.find('MsgId').text) if root.find('MsgId') is not None else None, } return message_data async def run_task(self, host: str, port: int, *args, **kwargs): """ 启动 Quart 应用。 """ await self.app.run_task(host=host, port=port, *args, **kwargs) def on_message(self, msg_type: str): """ 注册消息类型处理器。 """ def decorator(func: Callable[[OAEvent], None]): if msg_type not in self._message_handlers: self._message_handlers[msg_type] = [] self._message_handlers[msg_type].append(func) return func return decorator async def _handle_message(self, event: OAEvent): """ 处理消息事件。 """ message_id = event.message_id self._prune_state() if message_id in self.msg_id_map.keys(): self.msg_id_map[message_id] += 1 self._msg_seen_at[message_id] = time.monotonic() return self.msg_id_map[message_id] = 1 self._msg_seen_at[message_id] = time.monotonic() msg_type = event.type if msg_type in self._message_handlers: for handler in self._message_handlers[msg_type]: await handler(event) async def set_message(self, msg_id: int, content: str): self.generated_content[msg_id] = str(content)[: self._MAX_CONTENT_CHARS] self._generated_at[msg_id] = time.monotonic() self._prune_state() class OAClientForLongerResponse: _MAX_USERS = 4096 _MAX_MESSAGES_PER_USER = 20 _MAX_CONTENT_CHARS = 200000 def __init__( self, token: str, EncodingAESKey: str, AppID: str, Appsecret: str, LoadingMessage: str, logger: None, unified_mode: bool = False, api_base_url: str = 'https://api.weixin.qq.com', ): self.token = token self.aes = EncodingAESKey self.appid = AppID self.appsecret = Appsecret self.base_url = api_base_url self.access_token = '' self.unified_mode = unified_mode self.app = Quart(__name__) self.app.config['MAX_CONTENT_LENGTH'] = _MAX_CALLBACK_BODY_BYTES # 只有在非统一模式下才注册独立路由 if not self.unified_mode: self.app.add_url_rule( '/callback/command', 'handle_callback', self.handle_callback_request, methods=['GET', 'POST'], ) self._message_handlers = { 'example': [], } self.access_token_expiry_time = None self.loading_message = LoadingMessage self.msg_queue = {} self.user_msg_queue = {} self._last_queue_cleanup = 0.0 self.logger = logger def _prune_queues(self) -> None: now = time.monotonic() if now - self._last_queue_cleanup >= 60: self._last_queue_cleanup = now for user_id, queue in tuple(self.msg_queue.items()): if not queue: self.msg_queue.pop(user_id, None) for user_id, queue in tuple(self.user_msg_queue.items()): if not queue: self.user_msg_queue.pop(user_id, None) while len(self.msg_queue) > self._MAX_USERS: self.msg_queue.pop(next(iter(self.msg_queue)), None) while len(self.user_msg_queue) > self._MAX_USERS: self.user_msg_queue.pop(next(iter(self.user_msg_queue)), None) def clear(self) -> None: self.msg_queue.clear() self.user_msg_queue.clear() async def handle_callback_request(self): """处理回调请求(独立端口模式,使用全局 request)。""" return await self._handle_callback_internal(request) async def handle_unified_webhook(self, req): """处理回调请求(统一 webhook 模式,显式传递 request)。 Args: req: Quart Request 对象 Returns: 响应数据 """ return await self._handle_callback_internal(req) async def _handle_callback_internal(self, req): """处理回调请求的内部实现,包括 GET 验证和 POST 消息接收。 Args: req: Quart Request 对象 """ try: signature = req.args.get('signature', '') timestamp = req.args.get('timestamp', '') nonce = req.args.get('nonce', '') echostr = req.args.get('echostr', '') msg_signature = req.args.get('msg_signature', '') if msg_signature is None: await self.logger.error('msg_signature不在请求体中') raise Exception('msg_signature不在请求体中') if req.method == 'GET': check_str = ''.join(sorted([self.token, timestamp, nonce])) check_signature = hashlib.sha1(check_str.encode('utf-8')).hexdigest() return echostr if check_signature == signature else '拒绝请求' elif req.method == 'POST': encryt_msg = await req.data if len(encryt_msg) > _MAX_CALLBACK_BODY_BYTES: raise ValueError('Official Account callback body exceeds the size limit') wxcpt = WXBizMsgCrypt(self.token, self.aes, self.appid) ret, xml_msg = await asyncio.to_thread( wxcpt.DecryptMsg, encryt_msg, msg_signature, timestamp, nonce, ) xml_msg = xml_msg.decode('utf-8') if ret != 0: await self.logger.error('消息解密失败') raise Exception('消息解密失败') # 解析 XML root = await asyncio.to_thread(ET.fromstring, xml_msg) from_user = root.find('FromUserName').text to_user = root.find('ToUserName').text if self.msg_queue.get(from_user) and self.msg_queue[from_user][0]['content']: queue_top = self.msg_queue[from_user].pop(0) queue_content = queue_top['content'] # 弹出用户消息 if self.user_msg_queue.get(from_user) and self.user_msg_queue[from_user]: self.user_msg_queue[from_user].pop(0) self._prune_queues() response_xml = xml_template.format( to_user=from_user, from_user=to_user, create_time=int(time.time()), content=queue_content, ) return response_xml else: response_xml = xml_template.format( to_user=from_user, from_user=to_user, create_time=int(time.time()), content=self.loading_message, ) if self.user_msg_queue.get(from_user) and self.user_msg_queue[from_user][0]['content']: return response_xml else: message_data = await self.get_message(xml_msg) if message_data: event = OAEvent.from_payload(message_data) if event: self.user_msg_queue.setdefault(from_user, []).append( { 'content': str(event.message)[: self._MAX_CONTENT_CHARS], } ) self.user_msg_queue[from_user] = self.user_msg_queue[from_user][ -self._MAX_MESSAGES_PER_USER : ] self._prune_queues() await self._handle_message(event) return response_xml except Exception: await self.logger.error(f'handle_callback_request失败: {traceback.format_exc()}') traceback.print_exc() async def get_message(self, xml_msg: str): root = await asyncio.to_thread(ET.fromstring, xml_msg) message_data = { 'ToUserName': root.find('ToUserName').text, 'FromUserName': root.find('FromUserName').text, 'CreateTime': int(root.find('CreateTime').text), 'MsgType': root.find('MsgType').text, 'Content': root.find('Content').text if root.find('Content') is not None else None, 'MsgId': int(root.find('MsgId').text) if root.find('MsgId') is not None else None, } return message_data async def run_task(self, host: str, port: int, *args, **kwargs): """ 启动 Quart 应用。 """ await self.app.run_task(host=host, port=port, *args, **kwargs) def on_message(self, msg_type: str): """ 注册消息类型处理器。 """ def decorator(func: Callable[[OAEvent], None]): if msg_type not in self._message_handlers: self._message_handlers[msg_type] = [] self._message_handlers[msg_type].append(func) return func return decorator async def _handle_message(self, event: OAEvent): """ 处理消息事件。 """ msg_type = event.type if msg_type in self._message_handlers: for handler in self._message_handlers[msg_type]: await handler(event) async def set_message(self, from_user: int, message_id: int, content: str): if from_user not in self.msg_queue: self.msg_queue[from_user] = [] self.msg_queue[from_user].append( { 'msg_id': message_id, 'content': str(content)[: self._MAX_CONTENT_CHARS], } ) self.msg_queue[from_user] = self.msg_queue[from_user][-self._MAX_MESSAGES_PER_USER :] self._prune_queues()