mirror of
https://github.com/zhenxun-org/zhenxun_bot.git
synced 2026-10-03 10:40:02 +08:00
refactor code
This commit is contained in:
+175
-103
@@ -1,104 +1,128 @@
|
||||
import time
|
||||
from datetime import datetime, timedelta
|
||||
from collections import defaultdict
|
||||
from nonebot import require
|
||||
import nonebot
|
||||
import json
|
||||
import pytz
|
||||
from configs.path_config import TXT_PATH
|
||||
from configs.config import system_proxy
|
||||
from configs.config import SYSTEM_PROXY
|
||||
from typing import List, Union
|
||||
from nonebot.adapters import Bot
|
||||
import nonebot
|
||||
import pytz
|
||||
import pypinyin
|
||||
import aiohttp
|
||||
import time
|
||||
|
||||
try:
|
||||
import ujson as json
|
||||
except ModuleNotFoundError:
|
||||
import json
|
||||
|
||||
|
||||
scheduler = require('nonebot_plugin_apscheduler').scheduler
|
||||
scheduler = require("nonebot_plugin_apscheduler").scheduler
|
||||
|
||||
|
||||
# 次数检测
|
||||
class CountLimiter:
|
||||
def __init__(self, max):
|
||||
self.count = defaultdict(int)
|
||||
self.max = max
|
||||
"""
|
||||
次数检测工具,检测调用次数是否超过设定值
|
||||
"""
|
||||
|
||||
def add(self, key):
|
||||
def __init__(self, max_count: int):
|
||||
self.count = defaultdict(int)
|
||||
self.max_count = max_count
|
||||
|
||||
def add(self, key: Union[str, int, float]):
|
||||
self.count[key] += 1
|
||||
|
||||
def check(self, key) -> bool:
|
||||
if self.count[key] >= self.max:
|
||||
def check(self, key: Union[str, int, float]) -> bool:
|
||||
if self.count[key] >= self.max_count:
|
||||
self.count[key] = 0
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
# 用户正在执行此命令
|
||||
class UserExistLimiter:
|
||||
"""
|
||||
检测用户是否正在调用命令
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
self.mbool = defaultdict(bool)
|
||||
self.flag_data = defaultdict(bool)
|
||||
self.time = time.time()
|
||||
|
||||
def set_True(self, key):
|
||||
def set_True(self, key: Union[str, int, float]):
|
||||
self.time = time.time()
|
||||
self.mbool[key] = True
|
||||
self.flag_data[key] = True
|
||||
|
||||
def set_False(self, key):
|
||||
self.mbool[key] = False
|
||||
def set_False(self, key: Union[str, int, float]):
|
||||
self.flag_data[key] = False
|
||||
|
||||
def check(self, key):
|
||||
def check(self, key: Union[str, int, float]) -> bool:
|
||||
if time.time() - self.time > 30:
|
||||
self.set_False(key)
|
||||
return False
|
||||
return self.mbool[key]
|
||||
return self.flag_data[key]
|
||||
|
||||
|
||||
# 命令cd
|
||||
class FreqLimiter:
|
||||
def __init__(self, default_cd_seconds):
|
||||
"""
|
||||
命令冷却,检测用户是否处于冷却状态
|
||||
"""
|
||||
|
||||
def __init__(self, default_cd_seconds: int):
|
||||
self.next_time = defaultdict(float)
|
||||
self.default_cd = default_cd_seconds
|
||||
|
||||
def check(self, key) -> bool:
|
||||
def check(self, key: Union[str, int, float]) -> bool:
|
||||
return time.time() >= self.next_time[key]
|
||||
|
||||
def start_cd(self, key, cd_time=0):
|
||||
self.next_time[key] = time.time() + (cd_time if cd_time > 0 else self.default_cd)
|
||||
def start_cd(self, key: Union[str, int, float], cd_time: int = 0):
|
||||
self.next_time[key] = time.time() + (
|
||||
cd_time if cd_time > 0 else self.default_cd
|
||||
)
|
||||
|
||||
def left_time(self, key) -> float:
|
||||
def left_time(self, key: Union[str, int, float]) -> float:
|
||||
return self.next_time[key] - time.time()
|
||||
|
||||
|
||||
static_flmt = FreqLimiter(15)
|
||||
|
||||
|
||||
# 恶意触发命令检测
|
||||
class BanCheckLimiter:
|
||||
"""
|
||||
恶意命令触发检测
|
||||
"""
|
||||
|
||||
def __init__(self, default_check_time: float = 5, default_count: int = 4):
|
||||
self.mint = defaultdict(int)
|
||||
self.mtime = defaultdict(float)
|
||||
self.default_check_time = default_check_time
|
||||
self.default_count = default_count
|
||||
|
||||
def add(self, key):
|
||||
def add(self, key: Union[str, int, float]):
|
||||
if self.mint[key] == 1:
|
||||
self.mtime[key] = time.time()
|
||||
self.mint[key] += 1
|
||||
|
||||
def check(self, key) -> bool:
|
||||
# print(self.mint[key])
|
||||
# print(time.time() - self.mtime[key])
|
||||
def check(self, key: Union[str, int, float]) -> bool:
|
||||
if time.time() - self.mtime[key] > self.default_check_time:
|
||||
self.mtime[key] = time.time()
|
||||
self.mint[key] = 0
|
||||
return False
|
||||
if self.mint[key] >= self.default_count and time.time() - self.mtime[key] < self.default_check_time:
|
||||
if (
|
||||
self.mint[key] >= self.default_count
|
||||
and time.time() - self.mtime[key] < self.default_check_time
|
||||
):
|
||||
self.mtime[key] = time.time()
|
||||
self.mint[key] = 0
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
# 每日次数
|
||||
class DailyNumberLimiter:
|
||||
tz = pytz.timezone('Asia/Shanghai')
|
||||
"""
|
||||
每日调用命令次数限制
|
||||
"""
|
||||
|
||||
tz = pytz.timezone("Asia/Shanghai")
|
||||
|
||||
def __init__(self, max_num):
|
||||
self.today = -1
|
||||
@@ -123,7 +147,13 @@ class DailyNumberLimiter:
|
||||
self.count[key] = 0
|
||||
|
||||
|
||||
def is_number(s) -> bool:
|
||||
def is_number(s: str) -> bool:
|
||||
"""
|
||||
说明:
|
||||
检测 s 是否为数字
|
||||
参数:
|
||||
:param s: 文本
|
||||
"""
|
||||
try:
|
||||
float(s)
|
||||
return True
|
||||
@@ -131,6 +161,7 @@ def is_number(s) -> bool:
|
||||
pass
|
||||
try:
|
||||
import unicodedata
|
||||
|
||||
unicodedata.numeric(s)
|
||||
return True
|
||||
except (TypeError, ValueError):
|
||||
@@ -138,115 +169,156 @@ def is_number(s) -> bool:
|
||||
return False
|
||||
|
||||
|
||||
# 获取bot
|
||||
def get_bot():
|
||||
def get_bot() -> Bot:
|
||||
"""
|
||||
说明:
|
||||
获取 bot 对象
|
||||
"""
|
||||
return list(nonebot.get_bots().values())[0]
|
||||
|
||||
|
||||
def get_message_at(data: str) -> list:
|
||||
def get_message_at(data: str) -> List[int]:
|
||||
"""
|
||||
说明:
|
||||
获取消息中所有的 at 对象的 qq
|
||||
参数:
|
||||
:param data: event.json()
|
||||
"""
|
||||
qq_list = []
|
||||
data = json.loads(data)
|
||||
try:
|
||||
for msg in data['message']:
|
||||
if msg['type'] == 'at':
|
||||
qq_list.append(int(msg['data']['qq']))
|
||||
return qq_list
|
||||
except Exception:
|
||||
return []
|
||||
for msg in data["message"]:
|
||||
if msg["type"] == "at":
|
||||
qq_list.append(int(msg["data"]["qq"]))
|
||||
return qq_list
|
||||
|
||||
|
||||
def get_message_imgs(data: str) -> list:
|
||||
def get_message_imgs(data: str) -> List[str]:
|
||||
"""
|
||||
说明:
|
||||
获取消息中所有的 图片 的链接
|
||||
参数:
|
||||
:param data: event.json()
|
||||
"""
|
||||
img_list = []
|
||||
data = json.loads(data)
|
||||
try:
|
||||
for msg in data['message']:
|
||||
if msg['type'] == 'image':
|
||||
img_list.append(msg['data']['url'])
|
||||
return img_list
|
||||
except Exception:
|
||||
return []
|
||||
for msg in data["message"]:
|
||||
if msg["type"] == "image":
|
||||
img_list.append(msg["data"]["url"])
|
||||
return img_list
|
||||
|
||||
|
||||
def get_message_text(data: str) -> str:
|
||||
"""
|
||||
说明:
|
||||
获取消息中 纯文本 的信息
|
||||
参数:
|
||||
:param data: event.json()
|
||||
"""
|
||||
data = json.loads(data)
|
||||
result = ''
|
||||
try:
|
||||
for msg in data['message']:
|
||||
if msg['type'] == 'text':
|
||||
result += msg['data']['text'].strip() + ' '
|
||||
return result.strip()
|
||||
except Exception:
|
||||
return ''
|
||||
result = ""
|
||||
for msg in data["message"]:
|
||||
if msg["type"] == "text":
|
||||
result += msg["data"]["text"].strip() + " "
|
||||
return result.strip()
|
||||
|
||||
|
||||
def get_message_type(data: str) -> str:
|
||||
return json.loads(data)['message_type']
|
||||
|
||||
|
||||
def get_message_record(data: str) -> str:
|
||||
def get_message_record(data: str) -> List[str]:
|
||||
"""
|
||||
说明:
|
||||
获取消息中所有 语音 的链接
|
||||
参数:
|
||||
:param data: event.json()
|
||||
"""
|
||||
record_list = []
|
||||
data = json.loads(data)
|
||||
try:
|
||||
for msg in data['message']:
|
||||
if msg['type'] == 'record':
|
||||
return msg['data']['url']
|
||||
return ''
|
||||
except Exception:
|
||||
return ''
|
||||
for msg in data["message"]:
|
||||
if msg["type"] == "record":
|
||||
record_list.append(msg["data"]["url"])
|
||||
return record_list
|
||||
|
||||
|
||||
def get_message_json(data: str) -> dict:
|
||||
def get_message_json(data: str) -> List[dict]:
|
||||
"""
|
||||
说明:
|
||||
获取消息中所有 json
|
||||
参数:
|
||||
:param data: event.json()
|
||||
"""
|
||||
json_list = []
|
||||
data = json.loads(data)
|
||||
try:
|
||||
for msg in data['message']:
|
||||
if msg['type'] == 'json':
|
||||
return msg['data']
|
||||
return {}
|
||||
except Exception:
|
||||
return {}
|
||||
|
||||
|
||||
def add_to_16(value):
|
||||
while len(value) % 16 != 0:
|
||||
value += '\0'
|
||||
return str.encode(value)
|
||||
for msg in data["message"]:
|
||||
if msg["type"] == "json":
|
||||
json_list.append(msg["data"])
|
||||
return json_list
|
||||
|
||||
|
||||
# 获取文本加密后的cookie
|
||||
def get_cookie_text(cookie_name: str) -> str:
|
||||
with open(TXT_PATH + "cookie/" + cookie_name + ".txt", 'r') as f:
|
||||
"""
|
||||
说明:
|
||||
获取 txt/cookie 目录下指定 cookie 的内容
|
||||
参数:
|
||||
:param cookie_name: cookie文件名称
|
||||
"""
|
||||
with open(TXT_PATH + "cookie/" + cookie_name + ".txt", "r") as f:
|
||||
return f.read()
|
||||
|
||||
|
||||
# 获取本地http代理
|
||||
def get_local_proxy():
|
||||
return system_proxy if system_proxy else None
|
||||
"""
|
||||
说明:
|
||||
获取 config.py 中设置的代理
|
||||
"""
|
||||
return SYSTEM_PROXY if SYSTEM_PROXY else None
|
||||
|
||||
|
||||
# 判断是否为中文
|
||||
def is_Chinese(word):
|
||||
def is_Chinese(word: str) -> bool:
|
||||
"""
|
||||
说明:
|
||||
判断字符串是否为纯中文
|
||||
参数:
|
||||
:param word: 文本
|
||||
"""
|
||||
for ch in word:
|
||||
if '\u4e00' <= ch <= '\u9fff':
|
||||
return True
|
||||
return False
|
||||
if not "\u4e00" <= ch <= "\u9fff":
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
async def user_avatar(qq):
|
||||
url = f'http://q1.qlogo.cn/g?b=qq&nk={qq}&s=160'
|
||||
async def user_avatar(qq: int) -> bytes:
|
||||
"""
|
||||
说明:
|
||||
快捷获取用户头像
|
||||
参数:
|
||||
:param qq: qq号
|
||||
"""
|
||||
url = f"http://q1.qlogo.cn/g?b=qq&nk={qq}&s=160"
|
||||
async with aiohttp.ClientSession() as session:
|
||||
async with session.get(url, proxy=get_local_proxy(), timeout=5) as response:
|
||||
return await response.read()
|
||||
|
||||
|
||||
async def group_avatar(group_id):
|
||||
url = f'http://p.qlogo.cn/gh/{group_id}/{group_id}/640/'
|
||||
async def group_avatar(group_id: int) -> bytes:
|
||||
"""
|
||||
说明:
|
||||
快捷获取用群头像
|
||||
参数:
|
||||
:param group_id: 群号
|
||||
"""
|
||||
url = f"http://p.qlogo.cn/gh/{group_id}/{group_id}/640/"
|
||||
async with aiohttp.ClientSession() as session:
|
||||
async with session.get(url, proxy=get_local_proxy(), timeout=5) as response:
|
||||
return await response.read()
|
||||
|
||||
|
||||
def cn2py(word) -> str:
|
||||
def cn2py(word: str) -> str:
|
||||
"""
|
||||
说明:
|
||||
将字符串转化为拼音
|
||||
参数:
|
||||
:param word: 文本
|
||||
"""
|
||||
temp = ""
|
||||
for i in pypinyin.pinyin(word, style=pypinyin.NORMAL):
|
||||
temp += ''.join(i)
|
||||
temp += "".join(i)
|
||||
return temp
|
||||
|
||||
|
||||
Reference in New Issue
Block a user