refactor code

This commit is contained in:
hibiki
2021-07-30 21:21:51 +08:00
parent 2ad891aa1e
commit cc24822dca
165 changed files with 7815 additions and 8174 deletions
+90 -46
View File
@@ -1,10 +1,9 @@
from services.db_context import db
from typing import Optional, List
class BagUser(db.Model):
__tablename__ = 'bag_users'
__tablename__ = "bag_users"
id = db.Column(db.Integer(), primary_key=True)
user_qq = db.Column(db.BigInteger(), nullable=False)
belonging_group = db.Column(db.BigInteger(), nullable=False)
@@ -15,25 +14,41 @@ class BagUser(db.Model):
get_today_gold = db.Column(db.Integer(), default=0)
spend_today_gold = db.Column(db.Integer(), default=0)
_idx1 = db.Index('bag_group_users_idx1', 'user_qq', 'belonging_group', unique=True)
_idx1 = db.Index("bag_group_users_idx1", "user_qq", "belonging_group", unique=True)
@classmethod
async def get_my_total_gold(cls, user_qq: int, belonging_group: int) -> str:
"""
说明:
获取金币概况
参数:
:param user_qq: qq号
:param belonging_group: 所在群号
"""
query = cls.query.where(
(cls.user_qq == user_qq) & (cls.belonging_group == belonging_group)
)
user = await query.gino.first()
if not user:
user = await cls.create(
user_qq=user_qq,
belonging_group=belonging_group,
)
return f'当前金币:{user.gold}\n今日获取金币:{user.get_today_gold}\n今日花费金币:{user.spend_today_gold}' \
f'\n今日收益:{user.get_today_gold - user.spend_today_gold}' \
f'\n总赚取金币:{user.get_total_gold}\n总花费金币:{user.spend_total_gold}'
user_qq=user_qq,
belonging_group=belonging_group,
)
return (
f"当前金币:{user.gold}\n今日获取金币:{user.get_today_gold}\n今日花费金币:{user.spend_today_gold}"
f"\n今日收益:{user.get_today_gold - user.spend_today_gold}"
f"\n总赚取金币:{user.get_total_gold}\n总花费金币:{user.spend_total_gold}"
)
@classmethod
async def get_gold(cls, user_qq: int, belonging_group: int) -> int:
"""
说明:
获取当前金币
参数:
:param user_qq: qq号
:param belonging_group: 所在群号
"""
query = cls.query.where(
(cls.user_qq == user_qq) & (cls.belonging_group == belonging_group)
)
@@ -42,13 +57,20 @@ class BagUser(db.Model):
return user.gold
else:
await cls.create(
user_qq=user_qq,
belonging_group=belonging_group,
)
user_qq=user_qq,
belonging_group=belonging_group,
)
return 100
@classmethod
async def get_props(cls, user_qq: int, belonging_group: int) -> str:
"""
说明:
获取当前道具
参数:
:param user_qq: qq号
:param belonging_group: 所在群号
"""
query = cls.query.where(
(cls.user_qq == user_qq) & (cls.belonging_group == belonging_group)
)
@@ -57,13 +79,21 @@ class BagUser(db.Model):
return user.props
else:
await cls.create(
user_qq=user_qq,
belonging_group=belonging_group,
)
return ''
user_qq=user_qq,
belonging_group=belonging_group,
)
return ""
@classmethod
async def add_gold(cls, user_qq: int, belonging_group: int, num: int) -> bool:
"""
说明:
增加金币
参数:
:param user_qq: qq号
:param belonging_group: 所在群号
:param num: 金币数量
"""
query = cls.query.where(
(cls.user_qq == user_qq) & (cls.belonging_group == belonging_group)
)
@@ -74,7 +104,7 @@ class BagUser(db.Model):
await user.update(
gold=user.gold + num,
get_total_gold=user.get_total_gold + num,
get_today_gold=user.get_today_gold + num
get_today_gold=user.get_today_gold + num,
).apply()
else:
await cls.create(
@@ -90,6 +120,14 @@ class BagUser(db.Model):
@classmethod
async def spend_gold(cls, user_qq: int, belonging_group: int, num: int) -> bool:
"""
说明:
花费金币
参数:
:param user_qq: qq号
:param belonging_group: 所在群号
:param num: 金币数量
"""
query = cls.query.where(
(cls.user_qq == user_qq) & (cls.belonging_group == belonging_group)
)
@@ -100,7 +138,7 @@ class BagUser(db.Model):
await user.update(
gold=user.gold - num,
spend_total_gold=user.spend_total_gold + num,
spend_today_gold=user.spend_today_gold + num
spend_today_gold=user.spend_today_gold + num,
).apply()
else:
await cls.create(
@@ -108,7 +146,7 @@ class BagUser(db.Model):
belonging_group=belonging_group,
gold=100 - num,
spend_total_gold=num,
spend_today_gold=num
spend_today_gold=num,
)
return True
except Exception:
@@ -116,6 +154,14 @@ class BagUser(db.Model):
@classmethod
async def add_props(cls, user_qq: int, belonging_group: int, name: str) -> bool:
"""
说明:
增加道具
参数:
:param user_qq: qq号
:param belonging_group: 所在群号
:param name: 道具名称
"""
query = cls.query.where(
(cls.user_qq == user_qq) & (cls.belonging_group == belonging_group)
)
@@ -123,14 +169,10 @@ class BagUser(db.Model):
user = await query.gino.first()
try:
if user:
await user.update(
props=user.props + f'{name},'
).apply()
await user.update(props=user.props + f"{name},").apply()
else:
await cls.create(
user_qq=user_qq,
belonging_group=belonging_group,
props=f'{name},'
user_qq=user_qq, belonging_group=belonging_group, props=f"{name},"
)
return True
except Exception:
@@ -138,6 +180,14 @@ class BagUser(db.Model):
@classmethod
async def del_props(cls, user_qq: int, belonging_group: int, name: str) -> bool:
"""
说明:
使用道具
参数:
:param user_qq: qq号
:param belonging_group: 所在群号
:param name: 道具名称
"""
query = cls.query.where(
(cls.user_qq == user_qq) & (cls.belonging_group == belonging_group)
)
@@ -145,21 +195,19 @@ class BagUser(db.Model):
user = await query.gino.first()
try:
if user:
rst = ''
rst = ""
props = user.props
if props.find(name) != -1:
props = props.split(',')
props = props.split(",")
try:
index = props.index(name)
except ValueError:
return False
props = props[:index] + props[index + 1:]
props = props[:index] + props[index + 1 :]
for p in props:
if p != '':
rst += p + ','
await user.update(
props=rst
).apply()
if p != "":
rst += p + ","
await user.update(props=rst).apply()
return True
else:
return False
@@ -169,19 +217,15 @@ class BagUser(db.Model):
return False
@classmethod
async def get_user_all(cls, group_id: int = None) -> list:
async def get_all_users(cls, group_id: Optional[int] = None) -> List["BagUser"]:
"""
说明:
获取所有用户数据
参数:
:param group_id: 群号
"""
if not group_id:
query = await cls.query.gino.all()
else:
query = await cls.query.where(
(cls.belonging_group == group_id)
).gino.all()
query = await cls.query.where((cls.belonging_group == group_id)).gino.all()
return query
+48 -25
View File
@@ -3,20 +3,25 @@ import time
class BanUser(db.Model):
__tablename__ = 'ban_users'
__tablename__ = "ban_users"
user_qq = db.Column(db.BigInteger(), nullable=False, primary_key=True)
ban_level = db.Column(db.Integer(), nullable=False)
ban_time = db.Column(db.BigInteger())
duration = db.Column(db.BigInteger())
_idx1 = db.Index('ban_group_users_idx1', 'user_qq', unique=True)
_idx1 = db.Index("ban_group_users_idx1", "user_qq", unique=True)
@classmethod
async def check_ban_level(cls, user_qq: int, level: int) -> 'bool':
user = await cls.query.where(
(cls.user_qq == user_qq)
).gino.first()
async def check_ban_level(cls, user_qq: int, level: int) -> bool:
"""
说明:
检测ban掉目标的用户与unban用户的权限等级大小
参数:
:param user_qq: unban用户的qq号
:param level: ban掉目标用户的权限等级
"""
user = await cls.query.where((cls.user_qq == user_qq)).gino.first()
if not user:
return False
if user.ban_level > level:
@@ -24,21 +29,31 @@ class BanUser(db.Model):
return False
@classmethod
async def check_ban_time(cls, user_qq: int) -> 'str':
query = cls.query.where(
(cls.user_qq == user_qq)
)
async def check_ban_time(cls, user_qq: int) -> str:
"""
说明:
检测用户被ban时长
参数:
:param user_qq: qq号
"""
query = cls.query.where((cls.user_qq == user_qq))
user = await query.gino.first()
if not user:
return ''
return ""
if time.time() - (user.ban_time + user.duration) > 0 and user.duration != -1:
return ''
return ""
if user.duration == -1:
return '∞'
return "∞"
return time.time() - user.ban_time - user.duration
@classmethod
async def isban(cls, user_qq: int) -> 'bool':
async def is_ban(cls, user_qq: int) -> bool:
"""
说明:
判断用户是否被ban
参数:
:param user_qq: qq号
"""
if await cls.check_ban_time(user_qq):
return True
else:
@@ -46,10 +61,16 @@ class BanUser(db.Model):
return False
@classmethod
async def ban(cls, user_qq: int, ban_level: int, duration: int) -> 'bool':
query = cls.query.where(
(cls.user_qq == user_qq)
)
async def ban(cls, user_qq: int, ban_level: int, duration: int) -> bool:
"""
说明:
ban掉目标用户
参数:
:param user_qq: 目标用户qq号
:param ban_level: 使用ban命令用户的权限
:param duration: ban时长
"""
query = cls.query.where((cls.user_qq == user_qq))
query = query.with_for_update()
user = await query.gino.first()
if not await cls.check_ban_time(user_qq):
@@ -67,16 +88,18 @@ class BanUser(db.Model):
return False
@classmethod
async def unban(cls, user_qq: int) -> 'bool':
query = cls.query.where(
(cls.user_qq == user_qq)
)
async def unban(cls, user_qq: int) -> bool:
"""
说明:
unban用户
参数:
:param user_qq: qq号
"""
query = cls.query.where((cls.user_qq == user_qq))
query = query.with_for_update()
user = await query.gino.first()
if user is None:
return False
else:
await cls.delete.where(
(cls.user_qq == user_qq)
).gino.status()
await cls.delete.where((cls.user_qq == user_qq)).gino.status()
return True
+1 -1
View File
@@ -1,4 +1,4 @@
from datetime import datetime
from datetime import datetime
from services.db_context import db
+41 -56
View File
@@ -1,78 +1,63 @@
from services.db_context import db
class UserCount(db.Model):
__tablename__ = 'count_users'
__tablename__ = "count_users"
user_qq = db.Column(db.BigInteger(), nullable=False, primary_key=True)
reimu_count = db.Column(db.Integer(), nullable=False, default=0)
setu_r18_count = db.Column(db.Integer(), nullable=False, default=0)
_idx1 = db.Index('sign_reimu_users_idx1', 'user_qq', unique=True)
@classmethod
async def add_user(cls, user_qq: int):
query = cls.query.where(
(cls.user_qq == user_qq)
)
query = query.with_for_update()
if not await query.gino.first():
await cls.create(
user_qq=user_qq,
)
_idx1 = db.Index("sign_reimu_users_idx1", "user_qq", unique=True)
@classmethod
async def add_count(cls, user_qq: int, name: str, count: int = 1):
query = cls.query.where(
(cls.user_qq == user_qq)
)
"""
说明:
用户添加次数
参数:
:param user_qq: qq号
:param name: 目标名称
:param count: 增加次数
"""
query = cls.query.where((cls.user_qq == user_qq))
query = query.with_for_update()
user = await query.gino.first()
if user:
if name == 'reimu':
await user.update(
reimu_count=cls.reimu_count + count
).apply()
if name == 'setu_r18':
await user.update(
setu_r18_count=cls.setu_r18_count + count
).apply()
else:
await cls.create(
user_qq=user_qq
)
user = user if user else await cls.create(user_qq=user_qq)
if name == "reimu":
await user.update(reimu_count=cls.reimu_count + count).apply()
if name == "setu_r18":
await user.update(setu_r18_count=cls.setu_r18_count + count).apply()
@classmethod
async def check_count(cls, user_qq: int, name: str, max_count: int) -> bool:
query = cls.query.where(
(cls.user_qq == user_qq)
)
"""
说明:
检测次数是否到达最大值
参数:
:param user_qq: qq号
:param name: 目标名称
:param max_count: 最大值
"""
query = cls.query.where((cls.user_qq == user_qq))
user = await query.gino.first()
if user:
if name == 'reimu':
if user.reimu_count == max_count:
return True
else:
return False
if name == 'setu_r18':
if user.setu_r18_count == max_count:
return True
else:
return False
else:
await cls.add_user(user_qq)
return False
user = user if user else await cls.create(user_qq=user_qq)
if name == "reimu":
if user.reimu_count == max_count:
return True
else:
return False
if name == "setu_r18":
if user.setu_r18_count == max_count:
return True
else:
return False
@classmethod
async def reset_count(cls):
"""
说明:
重置每日次数
"""
for user in await cls.query.gino.all():
await user.update(
reimu_count=0,
setu_r18_count=0
).apply()
await user.update(reimu_count=0, setu_r18_count=0).apply()
+45 -23
View File
@@ -2,32 +2,41 @@ from services.db_context import db
class FriendUser(db.Model):
__tablename__ = 'friend_users'
__tablename__ = "friend_users"
id = db.Column(db.Integer(), primary_key=True)
user_id = db.Column(db.BigInteger(), nullable=False)
user_name = db.Column(db.Unicode(), nullable=False, default="")
nickname = db.Column(db.Unicode())
_idx1 = db.Index('friend_users_idx1', 'user_id', unique=True)
_idx1 = db.Index("friend_users_idx1", "user_id", unique=True)
@classmethod
async def get_user_name(cls, user_id: int) -> str:
query = cls.query.where(
cls.user_id == user_id
)
"""
说明:
获取好友用户名称
参数:
:param user_id: qq号
"""
query = cls.query.where(cls.user_id == user_id)
user = await query.gino.first()
if user:
return user.user_name
else:
return ''
return ""
@classmethod
async def add_friend_info(cls, user_id: int, user_name: str) -> 'bool':
async def add_friend_info(cls, user_id: int, user_name: str) -> bool:
"""
说明:
添加好友信息
参数:
:param user_id: qq号
:param user_name: 用户名称
"""
try:
query = cls.query.where(
cls.user_id == user_id
)
query = cls.query.where(cls.user_id == user_id)
user = await query.with_for_update().gino.first()
if not user:
await cls.create(
@@ -43,11 +52,15 @@ class FriendUser(db.Model):
return False
@classmethod
async def delete_friend_info(cls, user_id: int) -> 'bool':
async def delete_friend_info(cls, user_id: int) -> bool:
"""
说明:
删除好友信息
参数:
:param user_id: qq号
"""
try:
query = cls.query.where(
cls.user_id == user_id
)
query = cls.query.where(cls.user_id == user_id)
user = await query.with_for_update().gino.first()
if user:
await user.delete()
@@ -56,22 +69,31 @@ class FriendUser(db.Model):
return False
@classmethod
async def get_friend_nickname(cls, user_id: int) -> 'str':
query = cls.query.where(
cls.user_id == user_id
)
async def get_friend_nickname(cls, user_id: int) -> str:
"""
说明:
获取用户昵称
参数:
:param user_id: qq号
"""
query = cls.query.where(cls.user_id == user_id)
user = await query.gino.first()
if user:
if user.nickname:
return user.nickname
return ''
return ""
@classmethod
async def set_friend_nickname(cls, user_id: int, nickname: str) -> 'bool':
async def set_friend_nickname(cls, user_id: int, nickname: str) -> bool:
"""
说明:
设置用户昵称
参数:
:param user_id: qq号
:param nickname: 昵称
"""
try:
query = cls.query.where(
cls.user_id == user_id
)
query = cls.query.where(cls.user_id == user_id)
user = await query.with_for_update().gino.first()
if not user:
await cls.create(
+89 -58
View File
@@ -1,83 +1,124 @@
from services.db_context import db
from typing import Optional, List
class GoodsInfo(db.Model):
__tablename__ = 'goods_info'
__tablename__ = "goods_info"
id = db.Column(db.Integer(), primary_key=True)
goods_name = db.Column(db.TEXT(), nullable=False) # 名称
goods_price = db.Column(db.Integer(), nullable=False) # 价格
goods_description = db.Column(db.TEXT(), nullable=False) # 商品描述
goods_name = db.Column(db.TEXT(), nullable=False) # 名称
goods_price = db.Column(db.Integer(), nullable=False) # 价格
goods_description = db.Column(db.TEXT(), nullable=False) # 商品描述
goods_discount = db.Column(db.Numeric(scale=3, asdecimal=False), default=1) # 打折
goods_limit_time = db.Column(db.BigInteger(), default=0) # 限时
goods_limit_time = db.Column(db.BigInteger(), default=0) # 限时
_idx1 = db.Index('goods_group_users_idx1', 'goods_name', unique=True)
_idx1 = db.Index("goods_group_users_idx1", "goods_name", unique=True)
@classmethod
async def add_goods(cls, goods_name: str, goods_price: int,
goods_description: str, goods_discount: float = 1, goods_limit_time: int = 0) -> bool:
# try:
await cls.create(
goods_name=goods_name,
goods_price=goods_price,
goods_description=goods_description,
goods_discount=goods_discount,
goods_limit_time=goods_limit_time
async def add_goods(
cls,
goods_name: str,
goods_price: int,
goods_description: str,
goods_discount: float = 1,
goods_limit_time: int = 0,
) -> bool:
"""
说明:
添加商品
参数:
:param goods_name: 商品名称
:param goods_price: 商品价格
:param goods_description: 商品简介
:param goods_discount: 商品折扣
:param goods_limit_time: 商品限时
"""
try:
await cls.create(
goods_name=goods_name,
goods_price=goods_price,
goods_description=goods_description,
goods_discount=goods_discount,
goods_limit_time=goods_limit_time,
)
return True
except Exception:
return False
@classmethod
async def delete_goods(cls, goods_name: str) -> bool:
"""
说明:
删除商品
参数:
:param goods_name: 商品名称
"""
query = (
await cls.query.where(cls.goods_name == goods_name)
.with_for_update()
.gino.first()
)
return True
# except Exception:
# return False
@classmethod
async def del_goods(cls, goods_name: str) -> bool:
query = await cls.query.where(
cls.goods_name == goods_name
).with_for_update().gino.first()
if not query:
return False
await query.delete()
return True
@classmethod
async def update_goods(cls, goods_name: str, goods_price: int = None,
goods_description: str = None, goods_discount: float = None,
goods_limit_time: int = None) -> bool:
async def update_goods(
cls,
goods_name: str,
goods_price: Optional[int] = None,
goods_description: Optional[str] = None,
goods_discount: Optional[float] = None,
goods_limit_time: Optional[int] = None,
) -> bool:
"""
说明:
更新商品信息
参数:
:param goods_name: 商品名称
:param goods_price: 商品价格
:param goods_description: 商品简介
:param goods_discount: 商品折扣
:param goods_limit_time: 商品限时时间
"""
try:
query = await cls.query.where(
cls.goods_name == goods_name
).with_for_update().gino.first()
query = (
await cls.query.where(cls.goods_name == goods_name)
.with_for_update()
.gino.first()
)
if not query:
return False
if goods_price:
await query.update(
goods_price=goods_price
).apply()
await query.update(goods_price=goods_price).apply()
if goods_description:
await query.update(
goods_description=goods_description
).apply()
await query.update(goods_description=goods_description).apply()
if goods_discount:
await query.update(
goods_discount=goods_discount
).apply()
await query.update(goods_discount=goods_discount).apply()
if goods_limit_time:
await query.update(
goods_limit_time=goods_limit_time
).apply()
await query.update(goods_limit_time=goods_limit_time).apply()
return True
except Exception:
return False
@classmethod
async def get_goods_info(cls, goods_name: str) -> 'GoodsInfo':
query = await cls.query.where(
cls.goods_name == goods_name
).gino.first()
async def get_goods_info(cls, goods_name: str) -> "GoodsInfo":
"""
说明:
获取商品对象
参数:
:param goods_name: 商品名称
"""
query = await cls.query.where(cls.goods_name == goods_name).gino.first()
return query
@classmethod
async def get_all_goods(cls) -> list:
async def get_all_goods(cls) -> List["GoodsInfo"]:
"""
说明:
获得全部有序商品对象
"""
query = await cls.query.gino.all()
id_lst = [x.id for x in query]
goods_lst = []
@@ -86,13 +127,3 @@ class GoodsInfo(db.Model):
goods_lst.append([x for x in query if x.id == min_id][0])
id_lst.remove(min_id)
return goods_lst
+43 -15
View File
@@ -1,30 +1,47 @@
from services.db_context import db
from typing import List
class GroupInfo(db.Model):
__tablename__ = 'group_info'
__tablename__ = "group_info"
group_id = db.Column(db.BigInteger(), nullable=False, primary_key=True)
group_name = db.Column(db.Unicode(), nullable=False, default="")
max_member_count = db.Column(db.Integer(), nullable=False, default=0)
member_count = db.Column(db.Integer(), nullable=False, default=0)
_idx1 = db.Index('group_info_idx1', 'group_id', unique=True)
_idx1 = db.Index("group_info_idx1", "group_id", unique=True)
@classmethod
async def get_group_info(cls, group_id: int) -> 'GroupInfo':
query = cls.query.where(
cls.group_id == group_id
)
async def get_group_info(cls, group_id: int) -> "GroupInfo":
"""
说明:
获取群信息
参数:
:param group_id: 群号
"""
query = cls.query.where(cls.group_id == group_id)
return await query.gino.first()
@classmethod
async def add_group_info(cls, group_id: int, group_name: str, max_member_count: int, member_count: int) -> bool:
async def add_group_info(
cls, group_id: int, group_name: str, max_member_count: int, member_count: int
) -> bool:
"""
说明:
添加群信息
参数:
:param group_id: 群号
:param group_name: 群名称
:param max_member_count: 群员最大数量
:param member_count: 群员数量
"""
try:
group = await cls.query.where(
cls.group_id == group_id
).with_for_update().gino.first()
group = (
await cls.query.where(cls.group_id == group_id)
.with_for_update()
.gino.first()
)
if group:
await cls.update(
group_id=group_id,
@@ -45,12 +62,23 @@ class GroupInfo(db.Model):
@classmethod
async def delete_group_info(cls, group_id: int) -> bool:
"""
说明:
删除群信息
参数:
:param group_id: 群号
"""
try:
await cls.delete.where(
cls.group_id == group_id
).gino.status()
await cls.delete.where(cls.group_id == group_id).gino.status()
return True
except Exception:
return False
@classmethod
async def get_all_group(cls) -> List["GroupInfo"]:
"""
说明:
获取所有群对象
"""
query = await cls.query.gino.all()
return query
+69 -25
View File
@@ -1,10 +1,11 @@
from datetime import datetime
from services.db_context import db
from typing import List
class GroupInfoUser(db.Model):
__tablename__ = 'group_info_users'
__tablename__ = "group_info_users"
id = db.Column(db.Integer(), primary_key=True)
user_qq = db.Column(db.BigInteger(), nullable=False)
@@ -13,16 +14,30 @@ class GroupInfoUser(db.Model):
user_join_time = db.Column(db.DateTime(), nullable=False)
nickname = db.Column(db.Unicode())
_idx1 = db.Index('info_group_users_idx1', 'user_qq', 'belonging_group', unique=True)
_idx1 = db.Index("info_group_users_idx1", "user_qq", "belonging_group", unique=True)
@classmethod
async def insert(cls, user_qq: int, belonging_group: int, user_name: str, user_join_time: datetime) -> 'bool':
async def add_member_info(
cls,
user_qq: int,
belonging_group: int,
user_name: str,
user_join_time: datetime,
) -> bool:
"""
说明:
添加群内用户信息
参数:
:param user_qq: qq号
:param belonging_group: 群号
:param user_name: 用户名称
:param user_join_time: 入群时间
"""
query = cls.query.where(
(cls.user_qq == user_qq) & (cls.belonging_group == belonging_group)
)
query = query.with_for_update()
try:
if await query.gino.first() is None:
if not await query.gino.first():
await cls.create(
user_qq=user_qq,
user_name=user_name,
@@ -30,18 +45,34 @@ class GroupInfoUser(db.Model):
user_join_time=user_join_time,
)
return True
except:
except Exception:
return False
@classmethod
async def select_member_info(cls, user_qq: int, belonging_group: int) -> 'GroupInfoUser':
async def get_member_info(
cls, user_qq: int, belonging_group: int
) -> "GroupInfoUser":
"""
说明:
查询群员信息
参数:
:param user_qq: qq号
:param belonging_group: 群号
"""
query = cls.query.where(
(cls.user_qq == user_qq) & (cls.belonging_group == belonging_group)
)
return await query.gino.first()
@classmethod
async def delete_member_info(cls, user_qq: int, belonging_group: int) -> 'bool':
async def delete_member_info(cls, user_qq: int, belonging_group: int) -> bool:
"""
说明:
删除群员信息
参数:
:param user_qq: qq号
:param belonging_group: 群号
"""
query = cls.query.where(
(cls.user_qq == user_qq) & (cls.belonging_group == belonging_group)
)
@@ -55,34 +86,53 @@ class GroupInfoUser(db.Model):
(cls.user_qq == user_qq) & (cls.belonging_group == belonging_group)
).gino.status()
return True
except:
except Exception:
return False
@classmethod
async def query_group_member_list(cls, belonging_group: int) -> 'list':
async def get_group_member_id_list(cls, belonging_group: int) -> List[int]:
"""
说明:
获取该群所有用户qq
参数:
:param belonging_group: 群号
"""
member_list = []
query = cls.query.where(
(cls.belonging_group == belonging_group)
)
query = cls.query.where((cls.belonging_group == belonging_group))
for user in await query.gino.all():
member_list.append(user.user_qq)
return member_list
@classmethod
async def set_group_member_nickname(cls, user_qq: int, belonging_group: int, nickname: str) -> 'bool':
async def set_group_member_nickname(
cls, user_qq: int, belonging_group: int, nickname: str
) -> bool:
"""
说明:
设置群员在该群内的昵称
参数:
:param user_qq: qq号
:param belonging_group: 群号
:param nickname: 昵称
"""
query = cls.query.where(
(cls.user_qq == user_qq) & (cls.belonging_group == belonging_group)
)
user = await query.with_for_update().gino.first()
if user:
await user.update(
nickname=nickname
).apply()
await user.update(nickname=nickname).apply()
return True
return False
@classmethod
async def get_group_member_nickname(cls, user_qq: int, belonging_group: int) -> 'str':
async def get_group_member_nickname(cls, user_qq: int, belonging_group: int) -> str:
"""
说明:
获取用户在该群的昵称
参数:
:param user_qq: qq号
:param belonging_group: 群号
"""
query = cls.query.where(
(cls.user_qq == user_qq) & (cls.belonging_group == belonging_group)
)
@@ -90,10 +140,4 @@ class GroupInfoUser(db.Model):
if user:
if user.nickname:
return user.nickname
return ''
return ""
+47 -42
View File
@@ -2,99 +2,104 @@ from services.db_context import db
class GroupRemind(db.Model):
__tablename__ = 'group_reminds'
__tablename__ = "group_reminds"
id = db.Column(db.Integer(), primary_key=True)
group_id = db.Column(db.BigInteger(), nullable=False)
hy = db.Column(db.Boolean(), default=False) # 进群欢迎
kxcz = db.Column(db.Boolean(), default=False) # 开箱重置
zwa = db.Column(db.Boolean(), default=False) # 早晚安
gb = db.Column(db.Boolean(), default=True) # 广播
blpar = db.Column(db.Boolean(), default=True) # bilibili转发解析
pa = db.Column(db.Boolean(), default=True) # 爬
epic = db.Column(db.Boolean(), default=False) # epic
almanac = db.Column(db.Boolean(), default=False) # 原神黄历
hy = db.Column(db.Boolean(), default=False) # 进群欢迎
kxcz = db.Column(db.Boolean(), default=False) # 开箱重置
zwa = db.Column(db.Boolean(), default=False) # 早晚安
gb = db.Column(db.Boolean(), default=True) # 广播
blpar = db.Column(db.Boolean(), default=True) # bilibili转发解析
pa = db.Column(db.Boolean(), default=True) # 爬
epic = db.Column(db.Boolean(), default=False) # epic
almanac = db.Column(db.Boolean(), default=False) # 原神黄历
_idx1 = db.Index('info_group_reminds_idx1', 'group_id', unique=True)
_idx1 = db.Index("info_group_reminds_idx1", "group_id", unique=True)
@classmethod
async def get_status(cls, group_id: int, name: str) -> bool:
group = await cls.query.where(
(cls.group_id == group_id)
).gino.first()
"""
说明:
获取群通知状态
参数:
:param group_id: 群号
:param name: 目标名称
"""
group = await cls.query.where((cls.group_id == group_id)).gino.first()
if not group:
group = await cls.create(
group_id=group_id,
)
if name == 'hy':
if name == "hy":
return group.hy
if name == 'kxcz':
if name == "kxcz":
return group.kxcz
if name == 'zwa':
if name == "zwa":
return group.zwa
if name == 'gb':
if name == "gb":
return group.gb
if name == 'blpar':
if name == "blpar":
return group.blpar
if name == 'epic':
if name == "epic":
return group.epic
if name == 'pa':
if name == "pa":
return group.pa
if name == 'almanac':
if name == "almanac":
return group.almanac
@classmethod
async def set_status(cls, group_id: int, name: str, status: bool) -> bool:
"""
说明:
设置群通知状态
参数:
:param group_id: 群号
:param name: 目标名称
:param status: 通知状态
"""
try:
group = await cls.query.where(
(cls.group_id == group_id)
).with_for_update().gino.first()
group = (
await cls.query.where((cls.group_id == group_id))
.with_for_update()
.gino.first()
)
if not group:
group = await cls.create(
group_id=group_id,
)
if name == 'hy':
if name == "hy":
await group.update(
hy=status,
).apply()
if name == 'kxcz':
if name == "kxcz":
await group.update(
kxcz=status,
).apply()
if name == 'zwa':
if name == "zwa":
await group.update(
zwa=status,
).apply()
if name == 'gb':
if name == "gb":
await group.update(
gb=status,
).apply()
if name == 'blpar':
if name == "blpar":
await group.update(
blpar=status,
).apply()
if name == 'epic':
if name == "epic":
await group.update(
epic=status,
).apply()
if name == 'pa':
if name == "pa":
await group.update(
pa=status,
).apply()
if name == 'almanac':
if name == "almanac":
await group.update(
almanac=status,
).apply()
return True
except Exception as e:
return False
+50 -19
View File
@@ -2,7 +2,7 @@ from services.db_context import db
class LevelUser(db.Model):
__tablename__ = 'level_users'
__tablename__ = "level_users"
id = db.Column(db.Integer(), primary_key=True)
user_qq = db.Column(db.BigInteger(), nullable=False)
@@ -10,13 +10,18 @@ class LevelUser(db.Model):
user_level = db.Column(db.BigInteger(), nullable=False)
group_flag = db.Column(db.Integer(), nullable=False, default=0)
_idx1 = db.Index('level_group_users_idx1', 'user_qq', 'group_id', unique=True)
_idx1 = db.Index("level_group_users_idx1", "user_qq", "group_id", unique=True)
@classmethod
async def get_user_level(cls, user_qq: int, group_id: int) -> int:
query = cls.query.where(
(cls.user_qq == user_qq) & (cls.group_id == group_id)
)
"""
说明:
获取用户在群内的等级
参数:
:param user_qq: qq号
:param group_id: 群号
"""
query = cls.query.where((cls.user_qq == user_qq) & (cls.group_id == group_id))
user = await query.gino.first()
if user:
return user.user_level
@@ -24,10 +29,19 @@ class LevelUser(db.Model):
return -1
@classmethod
async def set_level(cls, user_qq: int, group_id: int, level: int, group_flag: int = 0) -> 'bool':
query = cls.query.where(
(cls.user_qq == user_qq) & (cls.group_id == group_id)
)
async def set_level(
cls, user_qq: int, group_id: int, level: int, group_flag: int = 0
) -> bool:
"""
说明:
设置用户在群内的权限
参数:
:param user_qq: qq号
:param group_id: 群号
:param level: 权限等级
:param group_flag: 是否被自动更新刷新权限 0:是,1:否
"""
query = cls.query.where((cls.user_qq == user_qq) & (cls.group_id == group_id))
query = query.with_for_update()
user = await query.gino.first()
if user is None:
@@ -43,10 +57,15 @@ class LevelUser(db.Model):
return False
@classmethod
async def delete_level(cls, user_qq: int, group_id: int) -> 'bool':
query = cls.query.where(
(cls.user_qq == user_qq) & (cls.group_id == group_id)
)
async def delete_level(cls, user_qq: int, group_id: int) -> bool:
"""
说明:
删除用户权限
参数:
:param user_qq: qq号
:param group_id: 群号
"""
query = cls.query.where((cls.user_qq == user_qq) & (cls.group_id == group_id))
query = query.with_for_update()
user = await query.gino.first()
if user is None:
@@ -56,7 +75,15 @@ class LevelUser(db.Model):
return True
@classmethod
async def check_level(cls, user_qq: int, group_id: int, level: int) -> 'bool':
async def check_level(cls, user_qq: int, group_id: int, level: int) -> bool:
"""
说明:
检查用户权限等级是否大于 level
参数:
:param user_qq: qq号
:param group_id: 群号
:param level: 权限等级
"""
if group_id != 0:
query = cls.query.where(
(cls.user_qq == user_qq) & (cls.group_id == group_id)
@@ -66,9 +93,7 @@ class LevelUser(db.Model):
return False
user_level = user.user_level
else:
query = cls.query.where(
cls.user_qq == user_qq
)
query = cls.query.where(cls.user_qq == user_qq)
highest_level = 0
for user in await query.gino.all():
if user.user_level > highest_level:
@@ -80,7 +105,14 @@ class LevelUser(db.Model):
return False
@classmethod
async def is_group_flag(cls, user_qq: int, group_id: int) -> 'bool':
async def is_group_flag(cls, user_qq: int, group_id: int) -> bool:
"""
说明:
检测是否会被自动更新刷新权限
参数:
:param user_qq: qq号
:param group_id: 群号
"""
user = await cls.query.where(
(cls.user_qq == user_qq) & (cls.group_id == group_id)
).gino.first()
@@ -90,4 +122,3 @@ class LevelUser(db.Model):
return True
else:
return False
+165
View File
@@ -0,0 +1,165 @@
from typing import Optional, List
from services.db_context import db
class Pixiv(db.Model):
__tablename__ = "pixiv"
id = db.Column(db.Integer(), primary_key=True)
pid = db.Column(db.BigInteger(), nullable=False)
title = db.Column(db.String(), nullable=False)
width = db.Column(db.Integer(), nullable=False)
height = db.Column(db.Integer(), nullable=False)
view = db.Column(db.Integer(), nullable=False)
bookmarks = db.Column(db.Integer(), nullable=False)
img_url = db.Column(db.String(), nullable=False)
img_p = db.Column(db.String(), nullable=False)
uid = db.Column(db.BigInteger(), nullable=False)
author = db.Column(db.String(), nullable=False)
is_r18 = db.Column(db.Boolean(), nullable=False)
tags = db.Column(db.String(), nullable=False)
_idx1 = db.Index("pixiv_idx1", "pid", "img_url", unique=True)
@classmethod
async def add_image_data(
cls,
pid: int,
title: str,
width: int,
height: int,
view: int,
bookmarks: int,
img_url: str,
img_p: str,
uid: int,
author: str,
tags: str,
):
"""
说明:
添加图片信息
参数:
:param pid: pid
:param title: 标题
:param width: 宽度
:param height: 长度
:param view: 被查看次数
:param bookmarks: 收藏数
:param img_url: url链接
:param img_p: 张数
:param uid: 作者uid
:param author: 作者名称
:param tags: 相关tag
"""
if not await cls.check_exists(pid, img_p):
await cls.create(
pid=pid,
title=title,
width=width,
height=height,
view=view,
bookmarks=bookmarks,
img_url=img_url,
img_p=img_p,
uid=uid,
author=author,
is_r18=True if 'R-18' in tags else False,
tags=tags,
)
return True
return False
@classmethod
async def remove_image_data(cls, pid: int, img_p: str) -> bool:
"""
说明:
删除图片数据
参数:
:param pid: 图片pid
:param img_p: 图片pid的张数,如:p0,p1
"""
try:
if img_p:
await cls.delete.where(
(cls.pid == pid) & (cls.img_p == img_p)
).gino.status()
else:
await cls.delete.where(cls.pid == pid).gino.status()
return True
except Exception:
return False
@classmethod
async def get_all_pid(cls) -> List[int]:
"""
说明:
获取所有PID
"""
pid = []
query = await cls.query.gino.all()
for image in query:
if image.pid not in pid:
pid.append(image.pid)
return pid
# 0:非r18 1:r18 2:混合
@classmethod
async def query_images(
cls,
keyword: Optional[List[str]] = None,
uid: Optional[int] = None,
pid: Optional[int] = None,
r18: int = 0,
) -> List[Optional["Pixiv"]]:
"""
说明:
查找符合条件的图片
参数:
:param keyword: 关键词
:param uid: 画师uid
:param pid: 图片pid
:param r18: 是否r18,0:非r18 1:r18 2:混合
"""
if r18 == 0:
query = await cls.query.where(cls.is_r18 == False).gino.all()
elif r18 == 1:
query = await cls.query.where(cls.is_r18 == True).gino.all()
else:
query = await cls.query.gino.all()
if keyword:
query = [x for x in query if set(x.tags.split(',')) > set(keyword)]
elif uid:
query = [x for x in query if x.uid == uid]
elif pid:
query = [x for x in query if x.pid == pid]
return query
@classmethod
async def check_exists(cls, pid: int, img_p: str) -> bool:
"""
说明:
检测pid是否已存在
参数:
:param pid: 图片PID
:param img_p: 张数
"""
query = await cls.query.where(
(cls.pid == pid) & (cls.img_p == img_p)
).gino.all()
return bool(query)
@classmethod
async def get_keyword_num(cls, keyword: List[str]) -> "int, int":
"""
说明:
获取相关关键词(keyword, tag)在图库中的数量
参数:
:param keyword: 关键词/Tag
"""
query = await cls.query.gino.all()
query = [x for x in query if set(x.tags.split(',')) > set(keyword)]
r18_count = len([x for x in query if x.is_r18])
return len(query) - r18_count, r18_count
+126
View File
@@ -0,0 +1,126 @@
from services.db_context import db
from typing import Set, List
class PixivKeywordUser(db.Model):
__tablename__ = "pixiv_keyword_users"
id = db.Column(db.Integer(), primary_key=True)
user_qq = db.Column(db.BigInteger(), nullable=False)
group_id = db.Column(db.BigInteger(), nullable=False)
keyword = db.Column(db.String(), nullable=False)
is_pass = db.Column(db.Boolean(), default=False)
_idx1 = db.Index("pixiv_keyword_users_idx1", "keyword", unique=True)
@classmethod
async def add_keyword(
cls, user_qq: int, group_id: int, keyword: str, superusers: Set[str]
) -> bool:
"""
说明:
添加搜图的关键词
参数:
:param user_qq: qq号
:param group_id: 群号
:param keyword: 关键词
:param superusers: 是否为超级用户
"""
is_pass = True if str(user_qq) in superusers else False
if not await cls._check_keyword_exists(keyword):
await cls.create(
user_qq=user_qq, group_id=group_id, keyword=keyword, is_pass=is_pass
)
return True
return False
@classmethod
async def delete_keyword(cls, keyword: str) -> bool:
"""
说明:
删除关键词
参数:
:param keyword: 关键词
"""
if await cls._check_keyword_exists(keyword):
query = cls.query.where(cls.keyword == keyword).with_for_update()
query = await query.gino.first()
await query.delete()
return True
return False
@classmethod
async def set_keyword_pass(cls, keyword: str, is_pass: bool) -> "int, int":
"""
说明:
通过或禁用关键词
参数:
:param keyword: 关键词
:param is_pass: 通过状态
"""
if await cls._check_keyword_exists(keyword):
query = cls.query.where(cls.keyword == keyword).with_for_update()
query = await query.gino.first()
await query.update(
is_pass=is_pass,
).apply()
return query.user_qq, query.group_id
return 0, 0
@classmethod
async def get_all_user_dict(cls) -> dict:
"""
说明:
获取关键词数据库各个用户贡献的关键词字典
"""
tmp = {}
query = await cls.query.gino.all()
for user in query:
if not tmp.get(user.user_qq):
tmp[user.user_qq] = {"keyword": []}
tmp[user.user_qq]["keyword"].append(user.keyword)
return tmp
@classmethod
async def get_current_keyword(cls) -> "List[str], List[str]":
"""
说明:
获取当前通过与未通过的关键词
"""
pass_keyword = []
not_pass_keyword = []
query = await cls.query.gino.all()
for user in query:
if user.is_pass:
pass_keyword.append(user.keyword)
else:
not_pass_keyword.append(user.keyword)
return pass_keyword, not_pass_keyword
@classmethod
async def get_black_pid(cls) -> List[str]:
"""
说明:
获取黑名单PID
"""
black_pid = []
query = await cls.query.where(cls.user_qq == 114514).gino.all()
for image in query:
black_pid.append(image.keyword[6:])
return black_pid
@classmethod
async def _check_keyword_exists(cls, keyword: str) -> bool:
"""
说明:
检测关键词是否已存在
参数:
:param keyword: 关键词
"""
current_keyword = []
query = await cls.query.gino.all()
for user in query:
current_keyword.append(user.keyword)
if keyword in current_keyword:
return True
return False
+33 -25
View File
@@ -1,9 +1,9 @@
from services.db_context import db
from typing import List
class RedbagUser(db.Model):
__tablename__ = 'redbag_users'
__tablename__ = "redbag_users"
id = db.Column(db.Integer(), primary_key=True)
user_qq = db.Column(db.BigInteger(), nullable=False)
@@ -13,18 +13,25 @@ class RedbagUser(db.Model):
spend_gold = db.Column(db.Integer(), default=0)
get_gold = db.Column(db.Integer(), default=0)
_idx1 = db.Index('redbag_group_users_idx1', 'user_qq', 'group_id', unique=True)
_idx1 = db.Index("redbag_group_users_idx1", "user_qq", "group_id", unique=True)
@classmethod
async def add_redbag_data(cls, user_qq: int, group_id: int, itype: str, money: int):
query = cls.query.where(
(cls.user_qq == user_qq) & (cls.group_id == group_id)
)
"""
说明:
添加收发红包数据
参数:
:param user_qq: qq号
:param group_id: 群号
:param itype: 收或发
:param money: 金钱数量
"""
query = cls.query.where((cls.user_qq == user_qq) & (cls.group_id == group_id))
user = await query.with_for_update().gino.first() or await cls.create(
user_qq=user_qq,
group_id=group_id,
)
if itype == 'get':
user_qq=user_qq,
group_id=group_id,
)
if itype == "get":
await user.update(
get_redbag_count=user.get_redbag_count + 1,
get_gold=user.get_gold + money,
@@ -37,9 +44,14 @@ class RedbagUser(db.Model):
@classmethod
async def ensure(cls, user_qq: int, group_id: int) -> bool:
query = cls.query.where(
(cls.user_qq == user_qq) & (cls.group_id == group_id)
)
"""
说明:
获取用户对象
参数:
:param user_qq: qq号
:param group_id: 群号
"""
query = cls.query.where((cls.user_qq == user_qq) & (cls.group_id == group_id))
user = await query.gino.first() or await cls.create(
user_qq=user_qq,
group_id=group_id,
@@ -47,19 +59,15 @@ class RedbagUser(db.Model):
return user
@classmethod
async def get_user_all(cls, group_id: int = None) -> list:
async def get_user_all(cls, group_id: int = None) -> List["RedbagUser"]:
"""
说明:
获取所有用户对象
参数:
:param group_id: 群号
"""
if not group_id:
query = await cls.query.gino.all()
else:
query = await cls.query.where(
(cls.group_id == group_id)
).gino.all()
query = await cls.query.where((cls.group_id == group_id)).gino.all()
return query
+60 -32
View File
@@ -3,7 +3,7 @@ from typing import List
class RussianUser(db.Model):
__tablename__ = 'russian_users'
__tablename__ = "russian_users"
id = db.Column(db.Integer(), primary_key=True)
user_qq = db.Column(db.BigInteger(), nullable=False)
@@ -13,34 +13,49 @@ class RussianUser(db.Model):
make_money = db.Column(db.Integer(), default=0)
lose_money = db.Column(db.Integer(), default=0)
_idx1 = db.Index('russian_group_users_idx1', 'user_qq', 'group_id', unique=True)
_idx1 = db.Index("russian_group_users_idx1", "user_qq", "group_id", unique=True)
@classmethod
async def ensure(cls, user_qq: int, group_id: int) -> 'RussianUser':
user = await cls.query.where(
(cls.user_qq == user_qq) & (cls.group_id == group_id)
).with_for_update().gino.first()
return user or await cls.create(
user_qq=user_qq,
group_id=group_id
async def ensure(cls, user_qq: int, group_id: int) -> "RussianUser":
"""
说明:
获取用户对象
参数:
:param user_qq: qq号
:param group_id: 群号
"""
user = (
await cls.query.where((cls.user_qq == user_qq) & (cls.group_id == group_id))
.with_for_update()
.gino.first()
)
return user or await cls.create(user_qq=user_qq, group_id=group_id)
@classmethod
async def add_count(cls, user_qq: int, group_id: int, itype: str) -> bool:
"""
说明:
添加用户输赢次数
说明:
:param user_qq: qq号
:param group_id: 群号
:param itype: 输或赢 'win' or 'lose'
"""
try:
user = await cls.query.where(
(cls.user_qq == user_qq) & (cls.group_id == group_id)
).with_for_update().gino.first()
if not user:
user = await cls.create(
user_qq=user_qq,
group_id=group_id
user = (
await cls.query.where(
(cls.user_qq == user_qq) & (cls.group_id == group_id)
)
if itype == 'win':
.with_for_update()
.gino.first()
)
if not user:
user = await cls.create(user_qq=user_qq, group_id=group_id)
if itype == "win":
await user.update(
win_count=user.win_count + 1,
).apply()
elif itype == 'lose':
elif itype == "lose":
await user.update(
fail_count=user.fail_count + 1,
).apply()
@@ -50,20 +65,30 @@ class RussianUser(db.Model):
@classmethod
async def money(cls, user_qq: int, group_id: int, itype: str, count: int) -> bool:
"""
说明:
添加用户输赢金钱
参数:
:param user_qq: qq号
:param group_id: 群号
:param itype: 输或赢 'win' or 'lose'
:param count: 金钱数量
"""
try:
user = await cls.query.where(
(cls.user_qq == user_qq) & (cls.group_id == group_id)
).with_for_update().gino.first()
if not user:
user = await cls.create(
user_qq=user_qq,
group_id=group_id
user = (
await cls.query.where(
(cls.user_qq == user_qq) & (cls.group_id == group_id)
)
if itype == 'win':
.with_for_update()
.gino.first()
)
if not user:
user = await cls.create(user_qq=user_qq, group_id=group_id)
if itype == "win":
await user.update(
make_money=user.make_money + count,
).apply()
elif itype == 'lose':
elif itype == "lose":
await user.update(
lose_money=user.lose_money + count,
).apply()
@@ -72,9 +97,12 @@ class RussianUser(db.Model):
return False
@classmethod
async def all_user(cls, group_id: int) -> List['RussianUser']:
users = await cls.query.where(
(cls.group_id == group_id)
).gino.all()
async def get_all_user(cls, group_id: int) -> List["RussianUser"]:
"""
说明:
获取该群所有用户对象
参数:
:param group_id: 群号
"""
users = await cls.query.where((cls.group_id == group_id)).gino.all()
return users
+168
View File
@@ -0,0 +1,168 @@
from services.db_context import db
from typing import List, Optional
class Setu(db.Model):
__tablename__ = "setu"
id = db.Column(db.Integer(), primary_key=True)
local_id = db.Column(db.Integer(), nullable=False)
title = db.Column(db.String(), nullable=False)
author = db.Column(db.String(), nullable=False)
pid = db.Column(db.BigInteger(), nullable=False)
img_hash = db.Column(db.String(), nullable=False)
img_url = db.Column(db.String(), nullable=False)
is_r18 = db.Column(db.Boolean(), nullable=False)
tags = db.Column(db.String())
_idx1 = db.Index("setu_pid_img_url_idx1", "pid", "img_url", unique=True)
@classmethod
async def add_setu_data(
cls,
local_id: int,
title: str,
author: str,
pid: int,
img_hash: str,
img_url: str,
tags: str,
):
"""
说明:
添加一份色图数据
参数:
:param local_id: 本地存储id
:param title: 标题
:param author: 作者
:param pid: 图片pid
:param img_hash: 图片hash值
:param img_url: 图片链接
:param tags: 图片标签
"""
if not await cls._check_exists(pid, img_url):
await cls.create(
local_id=local_id,
title=title,
author=author,
pid=pid,
img_hash=img_hash,
img_url=img_url,
is_r18=True if "R-18" in tags else False,
tags=tags,
)
@classmethod
async def query_image(
cls,
local_id: Optional[int] = None,
tags: Optional[List[str]] = None,
r18: int = 0,
):
"""
说明:
通过tag查找色图
参数:
:param local_id: 本地色图 id
:param tags: tags
:param r18: 是否 r18,0:非r18 1:r18 2:混合
"""
if local_id:
flag = True if r18 == 1 else False
return await cls.query.where(
(cls.local_id == local_id) & (cls.is_r18 == flag)
).gino.first()
if r18 == 0:
query = await cls.query.where(cls.is_r18 == False).gino.all()
elif r18 == 1:
query = await cls.query.where(cls.is_r18 == True).gino.all()
else:
query = await cls.query.gino.all()
if tags:
query = [x for x in query if set(x.tags.split(",")) > set(tags)]
return query
@classmethod
async def get_image_count(cls, r18: int = 0):
"""
说明:
查询图片数量
"""
return len(await cls.query_image(r18=r18))
# return await db.func.count(cls.local_id).gino.scalar()
@classmethod
async def get_image_in_hash(cls, img_hash: str) -> "Setu":
"""
说明:
通过图像hash获取图像信息
参数:
:param img_hash: = 图像hash值
"""
query = await cls.query.where(cls.img_hash == img_hash).gino.first()
return query
@classmethod
async def _check_exists(cls, pid: int, img_url: str) -> bool:
"""
说明:
检测图片是否存在
参数:
:param pid: 图片pid
:param img_url: 图片链接
"""
return bool(
await cls.query.where(
(cls.pid == pid) & (cls.img_url == img_url)
).gino.first()
)
@classmethod
async def update_setu_data(
cls,
pid: int,
*,
local_id: Optional[int] = None,
title: Optional[str] = None,
author: Optional[str] = None,
img_hash: Optional[str] = None,
img_url: Optional[str] = None,
tags: Optional[str] = None,
) -> bool:
"""
说明:
根据PID修改图片数据
参数:
:param local_id: 本地id
:param pid: 图片pid
:param title: 标题
:param author: 作者
:param img_hash: 图片hash值
:param img_url: 图片链接
:param tags: 图片标签
"""
query = cls.query.where(cls.pid == pid).with_for_update()
image_list = await query.gino.all()
if image_list:
for image in image_list:
if local_id:
await image.update(local_id=local_id).apply()
if title:
await image.update(title=title).apply()
if author:
await image.update(author=author).apply()
if img_hash:
await image.update(img_hash=img_hash).apply()
if img_url:
await image.update(img_url=img_url).apply()
if tags:
await image.update(tags=tags).apply()
return True
return False
@classmethod
async def get_all_setu(cls) -> List["Setu"]:
"""
说明:
获取所有图片对象
"""
return await cls.query.gino.all()
+19 -10
View File
@@ -4,7 +4,7 @@ from services.db_context import db
class SignGroupUser(db.Model):
__tablename__ = 'sign_group_users'
__tablename__ = "sign_group_users"
id = db.Column(db.Integer(), primary_key=True)
user_qq = db.Column(db.BigInteger(), nullable=False)
@@ -13,13 +13,19 @@ class SignGroupUser(db.Model):
checkin_count = db.Column(db.Integer(), nullable=False)
checkin_time_last = db.Column(db.DateTime(timezone=True), nullable=False)
impression = db.Column(db.Numeric(scale=3, asdecimal=False), nullable=False)
add_probability = db.Column(db.Numeric(scale=3, asdecimal=False), nullable=False, default=0)
specify_probability = db.Column(db.Numeric(scale=3, asdecimal=False), nullable=False, default=0)
add_probability = db.Column(
db.Numeric(scale=3, asdecimal=False), nullable=False, default=0
)
specify_probability = db.Column(
db.Numeric(scale=3, asdecimal=False), nullable=False, default=0
)
_idx1 = db.Index('sign_group_users_idx1', 'user_qq', 'belonging_group', unique=True)
_idx1 = db.Index("sign_group_users_idx1", "user_qq", "belonging_group", unique=True)
@classmethod
async def ensure(cls, user_qq: int, belonging_group: int, for_update: bool = False) -> 'SignGroupUser':
async def ensure(
cls, user_qq: int, belonging_group: int, for_update: bool = False
) -> "SignGroupUser":
query = cls.query.where(
(cls.user_qq == user_qq) & (cls.belonging_group == belonging_group)
)
@@ -35,14 +41,18 @@ class SignGroupUser(db.Model):
)
@classmethod
async def query_impression_all(cls, belonging_group: int) -> 'list,list':
async def get_all_impression(cls, belonging_group: int) -> "list, list":
"""
说明:
获取该群所有用户 id 及对应 好感度
参数:
:param belonging_group: 群号
"""
impression_list = []
user_qq_list = []
user_group = []
if belonging_group:
query = cls.query.where(
(cls.belonging_group == belonging_group)
)
query = cls.query.where(cls.belonging_group == belonging_group)
else:
query = cls.query
for user in await query.gino.all():
@@ -50,4 +60,3 @@ class SignGroupUser(db.Model):
user_qq_list.append(user.user_qq)
user_group.append(user.belonging_group)
return user_qq_list, impression_list, user_group