update v0.1.5.0

This commit is contained in:
HibiKier
2022-04-26 14:45:04 +08:00
parent c551e21766
commit 93539be492
22 changed files with 670 additions and 115 deletions
+102 -11
View File
@@ -1,5 +1,5 @@
from datetime import datetime, timedelta
from typing import List, Literal, Optional
from typing import List, Literal, Optional, Tuple, Union
from services.db_context import db
@@ -36,6 +36,86 @@ class ChatHistory(db.Model):
"""
return await cls._get_msg(uid, None, "user", msg_type, days).gino.all()
@classmethod
async def get_group_user_msg(
cls,
uid: int,
gid: int,
limit: int = 10,
date_scope: Tuple[datetime, datetime] = None,
) -> List["ChatHistory"]:
"""
说明:
获取群聊指定用户聊天记录
参数:
:param uid: qq
:param gid: 群号
:param limit: 获取数量
:param date_scope: 日期范围,默认None为全搜索
"""
return (
await cls._get_msg(uid, gid, "group", days=date_scope)
.limit(limit)
.gino.all()
)
@classmethod
async def get_group_user_msg_count(cls, uid: int, gid: int) -> Optional[int]:
"""
说明:
查询群聊指定用户的聊天记录数量
参数:
:param uid: qq
:param gid: 群号
"""
if x := await db.first(
db.text(
f"SELECT COUNT(id) as sum FROM public.chat_history WHERE user_qq = {uid} AND group_id = {gid}"
)
):
return x[0]
return None
@classmethod
async def get_group_msg_rank(
cls,
gid: int,
limit: int = 10,
order: str = "DESC",
date_scope: Optional[Tuple[datetime, datetime]] = None,
) -> Optional[Tuple[int, int]]:
"""
说明:
获取排行数据
参数:
:param gid: 群号
:param limit: 获取数量
:param order: 排序类型,desc,des
:param date_scope: 日期范围
"""
sql = f"SELECT user_qq, COUNT(id) as sum FROM public.chat_history WHERE group_id = {gid} "
if date_scope:
sql += f"AND create_time BETWEEN '{date_scope[0]}' AND '{date_scope[1]}' "
sql += f"GROUP BY user_qq ORDER BY sum {order if order and order.upper() != 'DES' else ''} LIMIT {limit}"
print(sql)
return await db.all(db.text(sql))
@classmethod
async def get_group_first_msg_datetime(cls, gid: int) -> Optional[datetime]:
"""
说明:
获取群第一条记录消息时间
参数:
:param gid:
"""
if (
msg := await cls.query.where(cls.group_id == gid)
.order_by(cls.create_time)
.gino.first()
):
return msg.create_time
return None
@classmethod
async def get_user_msg_count(
cls,
@@ -51,7 +131,9 @@ class ChatHistory(db.Model):
:param msg_type: 消息类型,私聊或群聊
:param days: 限制日期
"""
return (await cls._get_msg(uid, None, "user", msg_type, days, True).gino.first())[0]
return (
await cls._get_msg(uid, None, "user", msg_type, days, True).gino.first()
)[0]
@classmethod
async def get_group_msg(
@@ -81,7 +163,9 @@ class ChatHistory(db.Model):
:param gid: 用户qq
:param days: 限制日期
"""
return (await cls._get_msg(None, gid, "group", None, days, True).gino.first())[0]
return (await cls._get_msg(None, gid, "group", None, days, True).gino.first())[
0
]
@classmethod
def _get_msg(
@@ -89,9 +173,9 @@ class ChatHistory(db.Model):
uid: Optional[int],
gid: Optional[int],
type_: Literal["user", "group"],
msg_type: Optional[Literal["private", "group"]],
days: Optional[int],
is_select_count: bool = False
msg_type: Optional[Literal["private", "group"]] = None,
days: Optional[Union[int, Tuple[datetime, datetime]]] = None,
is_select_count: bool = False,
):
"""
说明:
@@ -104,8 +188,8 @@ class ChatHistory(db.Model):
:param days: 限制日期
"""
if is_select_count:
setattr(ChatHistory, 'count', db.func.count(cls.id).label('count'))
query = cls.select('count')
setattr(ChatHistory, "count", db.func.count(cls.id).label("count"))
query = cls.select("count")
else:
query = cls.query
if type_ == "user":
@@ -116,8 +200,15 @@ class ChatHistory(db.Model):
query = query.where(cls.group_id != None)
else:
query = query.where(cls.group_id == gid)
if uid:
query = query.where(cls.user_qq == uid)
if days:
query = query.where(
cls.create_time >= datetime.now() - timedelta(days=days)
)
if isinstance(days, int):
query = query.where(
cls.create_time >= datetime.now() - timedelta(days=days)
)
elif isinstance(days, tuple):
query = query.where(cls.create_time >= days[0]).where(
cls.create_time <= days[1]
)
return query