feat✨: 新增Web UI功能及数据库、日志等API接口

This commit is contained in:
HibiKier
2024-07-31 04:58:29 +08:00
parent f0b05ec5ed
commit 2bf5fd1a37
28 changed files with 2643 additions and 18 deletions
+1
View File
@@ -0,0 +1 @@
from .tabs import *
@@ -0,0 +1 @@
from .logs import *
@@ -0,0 +1,35 @@
import asyncio
from typing import Awaitable, Callable, Generic, TypeVar
PATTERN = r"\x1b(\[.*?[@-~]|\].*?(\x07|\x1b\\))"
_T = TypeVar("_T")
LogListener = Callable[[_T], Awaitable[None]]
class LogStorage(Generic[_T]):
"""
日志存储
"""
def __init__(self, rotation: float = 5 * 60):
self.count, self.rotation = 0, rotation
self.logs: dict[int, str] = {}
self.listeners: set[LogListener[str]] = set()
async def add(self, log: str):
seq = self.count = self.count + 1
self.logs[seq] = log
asyncio.get_running_loop().call_later(self.rotation, self.remove, seq)
await asyncio.gather(
*map(lambda listener: listener(log), self.listeners),
return_exceptions=True,
)
return seq
def remove(self, seq: int):
del self.logs[seq]
return
LOG_STORAGE: LogStorage[str] = LogStorage[str]()
+40
View File
@@ -0,0 +1,40 @@
from fastapi import APIRouter, WebSocket
from loguru import logger
from nonebot.utils import escape_tag
from starlette.websockets import WebSocket, WebSocketDisconnect, WebSocketState
from .log_manager import LOG_STORAGE
router = APIRouter()
@router.get("/logs", response_model=list[str])
async def system_logs_history(reverse: bool = False):
"""历史日志
参数:
reverse: 反转顺序.
"""
return LOG_STORAGE.list(reverse=reverse) # type: ignore
@router.websocket("/logs")
async def system_logs_realtime(websocket: WebSocket):
await websocket.accept()
async def log_listener(log: str):
await websocket.send_text(log)
LOG_STORAGE.listeners.add(log_listener)
try:
while websocket.client_state == WebSocketState.CONNECTED:
recv = await websocket.receive()
logger.trace(
f"{system_logs_realtime.__name__!r} received "
f"<e>{escape_tag(repr(recv))}</e>"
)
except WebSocketDisconnect:
pass
finally:
LOG_STORAGE.listeners.remove(log_listener)
return
@@ -0,0 +1,5 @@
from .database import *
from .main import *
from .manage import *
from .plugin_manage import *
from .system import *
@@ -0,0 +1,121 @@
import nonebot
from fastapi import APIRouter, Request
from nonebot.drivers import Driver
from tortoise import Tortoise
from tortoise.exceptions import OperationalError
from zhenxun.models.plugin_info import PluginInfo
from zhenxun.services.db_context import TestSQL
from ....base_model import BaseResultModel, QueryModel, Result
from ....utils import authentication
from .models.model import SqlModel, SqlText
from .models.sql_log import SqlLog
router = APIRouter(prefix="/database")
driver: Driver = nonebot.get_driver()
SQL_DICT = {}
SELECT_TABLE_SQL = """
select a.tablename as name,d.description as desc from pg_tables a
left join pg_class c on relname=tablename
left join pg_description d on oid=objoid and objsubid=0 where a.schemaname = 'public'
"""
SELECT_TABLE_COLUMN_SQL = """
SELECT column_name, data_type, character_maximum_length as max_length, is_nullable
FROM information_schema.columns
WHERE table_name = '{}';
"""
@driver.on_startup
async def _():
for plugin in nonebot.get_loaded_plugins():
module = plugin.name
sql_list = []
if plugin.metadata and plugin.metadata.extra:
sql_list = plugin.metadata.extra.get("sql_list")
if module in SQL_DICT:
raise ValueError(f"{module} 常用SQL module 重复")
if sql_list:
SqlModel(
name="",
module=module,
sql_list=sql_list,
)
SQL_DICT[module] = SqlModel
if SQL_DICT:
result = await PluginInfo.filter(module__in=SQL_DICT.keys()).values_list(
"module", "name"
)
module2name = {r[0]: r[1] for r in result}
for s in SQL_DICT:
module = SQL_DICT[s].module
if module in module2name:
SQL_DICT[s].name = module2name[module]
else:
SQL_DICT[s].name = module
@router.get(
"/get_table_list", dependencies=[authentication()], description="获取数据库表"
)
async def _() -> Result:
db = Tortoise.get_connection("default")
query = await db.execute_query_dict(SELECT_TABLE_SQL)
return Result.ok(query)
@router.get(
"/get_table_column", dependencies=[authentication()], description="获取表字段"
)
async def _(table_name: str) -> Result:
db = Tortoise.get_connection("default")
print(SELECT_TABLE_COLUMN_SQL.format(table_name))
query = await db.execute_query_dict(SELECT_TABLE_COLUMN_SQL.format(table_name))
return Result.ok(query)
@router.post("/exec_sql", dependencies=[authentication()], description="执行sql")
async def _(sql: SqlText, request: Request) -> Result:
ip = request.client.host if request.client else "unknown"
try:
if sql.sql.lower().startswith("select"):
db = Tortoise.get_connection("default")
res = await db.execute_query_dict(sql.sql)
await SqlLog.add(ip or "0.0.0.0", sql.sql, "")
return Result.ok(res, "执行成功啦!")
else:
result = await TestSQL.raw(sql.sql)
await SqlLog.add(ip or "0.0.0.0", sql.sql, str(result))
return Result.ok(info="执行成功啦!")
except OperationalError as e:
await SqlLog.add(ip or "0.0.0.0", sql.sql, str(e), False)
return Result.warning_(f"sql执行错误: {e}")
@router.post("/get_sql_log", dependencies=[authentication()], description="sql日志列表")
async def _(query: QueryModel) -> Result:
total = await SqlLog.all().count()
if total % query.size:
total += 1
data = (
await SqlLog.all()
.order_by("-id")
.offset((query.index - 1) * query.size)
.limit(query.size)
)
return Result.ok(BaseResultModel(total=total, data=data))
@router.get("/get_common_sql", dependencies=[authentication()], description="常用sql")
async def _(plugin_name: str | None = None) -> Result:
if plugin_name:
return Result.ok(SQL_DICT.get(plugin_name))
return Result.ok(str(SQL_DICT))
@@ -0,0 +1,24 @@
from pydantic import BaseModel
from zhenxun.utils.plugin_models.base import CommonSql
class SqlText(BaseModel):
"""
sql语句
"""
sql: str
class SqlModel(BaseModel):
"""
常用sql
"""
name: str
"""插件中文名称"""
module: str
"""插件名称"""
sql_list: list[CommonSql]
"""插件列表"""
@@ -0,0 +1,37 @@
from tortoise import fields
from zhenxun.services.db_context import Model
class SqlLog(Model):
id = fields.IntField(pk=True, generated=True, auto_increment=True)
"""自增id"""
ip = fields.CharField(255)
"""ip"""
sql = fields.CharField(255)
"""sql"""
result = fields.CharField(255, null=True)
"""结果"""
is_suc = fields.BooleanField(default=True)
"""是否成功"""
create_time = fields.DatetimeField(auto_now_add=True)
"""创建时间"""
class Meta:
table = "sql_log"
table_description = "sql执行日志"
@classmethod
async def add(
cls, ip: str, sql: str, result: str | None = None, is_suc: bool = True
):
"""获取用户在群内的等级
参数:
ip: ip
sql: sql
result: 返回结果
is_suc: 是否成功
"""
await cls.create(ip=ip, sql=sql, result=result, is_suc=is_suc)
@@ -0,0 +1,290 @@
import asyncio
import time
from datetime import datetime, timedelta
from pathlib import Path
import nonebot
from fastapi import APIRouter, WebSocket
from starlette.websockets import WebSocket, WebSocketDisconnect, WebSocketState
from tortoise.functions import Count
from websockets.exceptions import ConnectionClosedError, ConnectionClosedOK
from zhenxun.models.chat_history import ChatHistory
from zhenxun.models.group_info import GroupInfo
from zhenxun.models.plugin_info import PluginInfo
from zhenxun.models.statistics import Statistics
from zhenxun.services.log import logger
from zhenxun.utils.platform import PlatformUtils
from ....base_model import Result
from ....config import AVA_URL, GROUP_AVA_URL, QueryDateType
from ....utils import authentication, get_system_status
from .data_source import bot_live
from .model import ActiveGroup, BaseInfo, ChatHistoryCount, HotPlugin
run_time = time.time()
ws_router = APIRouter()
router = APIRouter(prefix="/main")
@router.get("/get_base_info", dependencies=[authentication()], description="基础信息")
async def _(bot_id: str | None = None) -> Result:
"""获取Bot基础信息
参数:
bot_id (Optional[str], optional): bot_id. Defaults to None.
返回:
Result: 获取指定bot信息与bot列表
"""
bot_list: list[BaseInfo] = []
if bots := nonebot.get_bots():
select_bot: BaseInfo
for key, bot in bots.items():
login_info = await bot.get_login_info()
bot_list.append(
BaseInfo(
bot=bot, # type: ignore
self_id=bot.self_id,
nickname=login_info["nickname"],
ava_url=AVA_URL.format(bot.self_id),
)
)
# 获取指定qq号的bot信息,若无指定 则获取第一个
if _bl := [b for b in bot_list if b.self_id == bot_id]:
select_bot = _bl[0]
else:
select_bot = bot_list[0]
select_bot.is_select = True
select_bot.config = select_bot.bot.config
now = datetime.now()
# 今日累计接收消息
select_bot.received_messages = await ChatHistory.filter(
bot_id=select_bot.self_id,
create_time__gte=now - timedelta(hours=now.hour),
).count()
# 群聊数量
select_bot.group_count = len(await select_bot.bot.get_group_list())
# 好友数量
select_bot.friend_count = len(await select_bot.bot.get_friend_list())
for bot in bot_list:
bot.bot = None # type: ignore
# 插件加载数量
select_bot.plugin_count = await PluginInfo.all().count()
fail_count = await PluginInfo.filter(load_status=False).count()
select_bot.fail_plugin_count = fail_count
select_bot.success_plugin_count = (
select_bot.plugin_count - select_bot.fail_plugin_count
)
# 连接时间
select_bot.connect_time = bot_live.get(select_bot.self_id) or 0
if select_bot.connect_time:
connect_date = datetime.fromtimestamp(select_bot.connect_time)
connect_date_str = connect_date.strftime("%Y-%m-%d %H:%M:%S")
select_bot.connect_date = datetime.strptime(
connect_date_str, "%Y-%m-%d %H:%M:%S"
)
version_file = Path() / "__version__"
if version_file.exists():
if text := version_file.open().read():
if ver := text.replace("__version__: ", "").strip():
select_bot.version = ver
day_call = await Statistics.filter(
create_time__gte=now - timedelta(hours=now.hour)
).count()
select_bot.day_call = day_call
return Result.ok(bot_list, "拿到信息啦!")
return Result.warning_("无Bot连接...")
@router.get(
"/get_all_ch_count", dependencies=[authentication()], description="获取接收消息数量"
)
async def _(bot_id: str) -> Result:
now = datetime.now()
all_count = await ChatHistory.filter(bot_id=bot_id).count()
day_count = await ChatHistory.filter(
bot_id=bot_id, create_time__gte=now - timedelta(hours=now.hour)
).count()
week_count = await ChatHistory.filter(
bot_id=bot_id, create_time__gte=now - timedelta(days=7)
).count()
month_count = await ChatHistory.filter(
bot_id=bot_id, create_time__gte=now - timedelta(days=30)
).count()
year_count = await ChatHistory.filter(
bot_id=bot_id, create_time__gte=now - timedelta(days=365)
).count()
return Result.ok(
ChatHistoryCount(
num=all_count,
day=day_count,
week=week_count,
month=month_count,
year=year_count,
)
)
@router.get(
"/get_ch_count", dependencies=[authentication()], description="获取接收消息数量"
)
async def _(bot_id: str, query_type: QueryDateType | None = None) -> Result:
if bots := nonebot.get_bots():
if not query_type:
return Result.ok(await ChatHistory.filter(bot_id=bot_id).count())
now = datetime.now()
if query_type == QueryDateType.DAY:
return Result.ok(
await ChatHistory.filter(
bot_id=bot_id, create_time__gte=now - timedelta(hours=now.hour)
).count()
)
if query_type == QueryDateType.WEEK:
return Result.ok(
await ChatHistory.filter(
bot_id=bot_id, create_time__gte=now - timedelta(days=7)
).count()
)
if query_type == QueryDateType.MONTH:
return Result.ok(
await ChatHistory.filter(
bot_id=bot_id, create_time__gte=now - timedelta(days=30)
).count()
)
if query_type == QueryDateType.YEAR:
return Result.ok(
await ChatHistory.filter(
bot_id=bot_id, create_time__gte=now - timedelta(days=365)
).count()
)
return Result.warning_("无Bot连接...")
@router.get(
"get_fg_count", dependencies=[authentication()], description="好友/群组数量"
)
async def _(bot_id: str) -> Result:
if bots := nonebot.get_bots():
if bot_id not in bots:
return Result.warning_("指定Bot未连接...")
bot = bots[bot_id]
platform = PlatformUtils.get_platform(bot)
if platform == "qq":
data = {
"friend_count": len(await bot.get_friend_list()),
"group_count": len(await bot.get_group_list()),
}
return Result.ok(data)
return Result.warning_("暂不支持该平台...")
return Result.warning_("无Bot连接...")
@router.get(
"/get_run_time", dependencies=[authentication()], description="获取nb运行时间"
)
async def _() -> Result:
return Result.ok(int(time.time() - run_time))
@router.get(
"/get_active_group", dependencies=[authentication()], description="获取活跃群聊"
)
async def _(date_type: QueryDateType | None = None) -> Result:
query = ChatHistory
now = datetime.now()
if date_type == QueryDateType.DAY:
query = ChatHistory.filter(create_time__gte=now - timedelta(hours=now.hour))
if date_type == QueryDateType.WEEK:
query = ChatHistory.filter(create_time__gte=now - timedelta(days=7))
if date_type == QueryDateType.MONTH:
query = ChatHistory.filter(create_time__gte=now - timedelta(days=30))
if date_type == QueryDateType.YEAR:
query = ChatHistory.filter(create_time__gte=now - timedelta(days=365))
data_list = (
await query.annotate(count=Count("id"))
.filter(group_id__not_isnull=True)
.group_by("group_id")
.order_by("-count")
.limit(5)
.values_list("group_id", "count")
)
active_group_list = []
id2name = {}
if data_list:
if info_list := await GroupInfo.filter(
group_id__in=[x[0] for x in data_list]
).all():
for group_info in info_list:
id2name[group_info.group_id] = group_info.group_name
for data in data_list:
active_group_list.append(
ActiveGroup(
group_id=data[0],
name=id2name.get(data[0]) or data[0],
chat_num=data[1],
ava_img=GROUP_AVA_URL.format(data[0], data[0]),
)
)
active_group_list = sorted(
active_group_list, key=lambda x: x.chat_num, reverse=True
)
if len(active_group_list) > 5:
active_group_list = active_group_list[:5]
return Result.ok(active_group_list)
@router.get(
"/get_hot_plugin", dependencies=[authentication()], description="获取热门插件"
)
async def _(date_type: QueryDateType | None = None) -> Result:
query = Statistics
now = datetime.now()
if date_type == QueryDateType.DAY:
query = Statistics.filter(create_time__gte=now - timedelta(hours=now.hour))
if date_type == QueryDateType.WEEK:
query = Statistics.filter(create_time__gte=now - timedelta(days=7))
if date_type == QueryDateType.MONTH:
query = Statistics.filter(create_time__gte=now - timedelta(days=30))
if date_type == QueryDateType.YEAR:
query = Statistics.filter(create_time__gte=now - timedelta(days=365))
data_list = (
await query.annotate(count=Count("id"))
.group_by("plugin_name")
.order_by("-count")
.limit(5)
.values_list("plugin_name", "count")
)
hot_plugin_list = []
module_list = [x[0] for x in data_list]
plugins = await PluginInfo.filter(module__in=module_list).all()
module2name = {p.module: p.name for p in plugins}
for data in data_list:
module = data[0]
name = module2name.get(module) or module
hot_plugin_list.append(
HotPlugin(
module=data[0],
name=name,
count=data[1],
)
)
hot_plugin_list = sorted(hot_plugin_list, key=lambda x: x.count, reverse=True)
if len(hot_plugin_list) > 5:
hot_plugin_list = hot_plugin_list[:5]
return Result.ok(hot_plugin_list)
@ws_router.websocket("/system_status")
async def system_logs_realtime(websocket: WebSocket, sleep: int = 5):
await websocket.accept()
logger.debug("ws system_status is connect")
try:
while websocket.client_state == WebSocketState.CONNECTED:
system_status = await get_system_status()
await websocket.send_text(system_status.json())
await asyncio.sleep(sleep)
except (WebSocketDisconnect, ConnectionClosedError, ConnectionClosedOK):
pass
return
@@ -0,0 +1,35 @@
import time
import nonebot
from nonebot.adapters.onebot.v11 import Bot
from nonebot.drivers import Driver
driver: Driver = nonebot.get_driver()
class BotLive:
def __init__(self):
self._data = {}
def add(self, bot_id: str):
self._data[bot_id] = time.time()
def get(self, bot_id: str) -> int | None:
return self._data.get(bot_id)
def remove(self, bot_id: str):
if bot_id in self._data:
del self._data[bot_id]
bot_live = BotLive()
@driver.on_bot_connect
async def _(bot: Bot):
bot_live.add(bot.self_id)
@driver.on_bot_disconnect
async def _(bot: Bot):
bot_live.remove(bot.self_id)
@@ -0,0 +1,105 @@
from datetime import datetime
from nonebot.adapters import Bot
from nonebot.config import Config
from pydantic import BaseModel
class SystemStatus(BaseModel):
"""
系统状态
"""
cpu: float
memory: float
disk: float
class BaseInfo(BaseModel):
"""
基础信息
"""
bot: Bot
"""Bot"""
self_id: str
"""SELF ID"""
nickname: str
"""昵称"""
ava_url: str
"""头像url"""
friend_count: int = 0
"""好友数量"""
group_count: int = 0
"""群聊数量"""
received_messages: int = 0
"""今日 累计接收消息"""
connect_time: int = 0
"""连接时间"""
connect_date: datetime | None = None
"""连接日期"""
plugin_count: int = 0
"""加载插件数量"""
success_plugin_count: int = 0
"""加载成功插件数量"""
fail_plugin_count: int = 0
"""加载失败插件数量"""
is_select: bool = False
"""当前选择"""
config: Config | None = None
"""nb配置"""
day_call: int = 0
"""今日调用插件次数"""
version: str = "unknown"
"""真寻版本"""
class Config:
arbitrary_types_allowed = True
class ChatHistoryCount(BaseModel):
"""
聊天记录数量
"""
num: int
"""总数"""
day: int
"""一天内"""
week: int
"""一周内"""
month: int
"""一月内"""
year: int
"""一年内"""
class ActiveGroup(BaseModel):
"""
活跃群聊数据
"""
group_id: str
"""群组id"""
name: str
"""群组名称"""
chat_num: int
"""发言数量"""
ava_img: str
"""群组头像"""
class HotPlugin(BaseModel):
"""
热门插件
"""
module: str
"""模块名"""
name: str
"""插件名称"""
count: int
"""调用次数"""
@@ -0,0 +1,529 @@
import re
from typing import Literal
import nonebot
from fastapi import APIRouter
from nonebot.adapters.onebot.v11 import ActionFailed
from starlette.websockets import WebSocket, WebSocketDisconnect, WebSocketState
from tortoise.functions import Count
from zhenxun.configs.config import NICKNAME
from zhenxun.models.ban_console import BanConsole
from zhenxun.models.chat_history import ChatHistory
from zhenxun.models.fg_request import FgRequest
from zhenxun.models.friend_user import FriendUser
from zhenxun.models.group_console import GroupConsole
from zhenxun.models.group_member_info import GroupInfoUser
from zhenxun.models.plugin_info import PluginInfo
from zhenxun.models.statistics import Statistics
from zhenxun.models.task_info import TaskInfo
from zhenxun.services.log import logger
from zhenxun.utils.enum import RequestHandleType, RequestType
from zhenxun.utils.exception import NotFoundError
from zhenxun.utils.platform import PlatformUtils
from ....base_model import Result
from ....config import AVA_URL, GROUP_AVA_URL
from ....utils import authentication
from ...logs.log_manager import LOG_STORAGE
from .model import (
DeleteFriend,
Friend,
FriendRequestResult,
GroupDetail,
GroupRequestResult,
GroupResult,
HandleRequest,
LeaveGroup,
Message,
MessageItem,
Plugin,
ReqResult,
SendMessage,
Task,
UpdateGroup,
UserDetail,
)
ws_router = APIRouter()
router = APIRouter(prefix="/manage")
SUB_PATTERN = r"\x1b(\[.*?[@-~]|\].*?(\x07|\x1b\\))"
GROUP_PATTERN = r'.*?Message (-?\d*) from (\d*)@\[群:(\d*)] "(.*)"'
PRIVATE_PATTERN = r'.*?Message (-?\d*) from (\d*) "(.*)"'
AT_PATTERN = r"\[CQ:at,qq=(.*)\]"
IMAGE_PATTERN = r"\[CQ:image,.*,url=(.*);.*?\]"
@router.get(
"/get_group_list", dependencies=[authentication()], description="获取群组列表"
)
async def _(bot_id: str) -> Result:
"""
获取群信息
"""
if bots := nonebot.get_bots():
if bot_id not in bots:
return Result.warning_("指定Bot未连接...")
group_list_result = []
try:
group_info = {}
group_list = await bots[bot_id].get_group_list()
for g in group_list:
gid = g["group_id"]
g["ava_url"] = GROUP_AVA_URL.format(gid, gid)
group_list_result.append(GroupResult(**g))
except Exception as e:
logger.error("调用API错误", "/get_group_list", e=e)
return Result.fail(f"{type(e)}: {e}")
return Result.ok(group_list_result, "拿到了新鲜出炉的数据!")
return Result.warning_("无Bot连接...")
@router.post(
"/update_group", dependencies=[authentication()], description="修改群组信息"
)
async def _(group: UpdateGroup) -> Result:
try:
group_id = group.group_id
if db_group := await GroupConsole.get_or_none(group_id=group_id):
db_group.level = group.level
db_group.status = group.status
if group.close_plugins:
db_group.block_plugin = ",".join(group.close_plugins) + ","
# TODO: 关闭task
await db_group.save(update_fields=["level", "status", "block_plugin"])
except Exception as e:
logger.error("调用API错误", "/get_group", e=e)
return Result.fail(f"{type(e)}: {e}")
return Result.ok(info="已完成记录!")
@router.get(
"/get_friend_list", dependencies=[authentication()], description="获取好友列表"
)
async def _(bot_id: str) -> Result:
"""
获取群信息
"""
if bots := nonebot.get_bots():
if bot_id not in bots:
return Result.warning_("指定Bot未连接...")
try:
platform = PlatformUtils.get_platform(bots[bot_id])
if platform != "qq":
return Result.warning_("该平台暂不支持该功能...")
friend_list = await bots[bot_id].get_friend_list()
for f in friend_list:
f["ava_url"] = AVA_URL.format(f["user_id"])
return Result.ok(
[Friend(**f) for f in friend_list if str(f["user_id"]) != bot_id],
"拿到了新鲜出炉的数据!",
)
except Exception as e:
logger.error("调用API错误", "/get_group_list", e=e)
return Result.fail(f"{type(e)}: {e}")
return Result.warning_("无Bot连接...")
@router.get(
"/get_request_count", dependencies=[authentication()], description="获取请求数量"
)
async def _() -> Result:
f_count = await FgRequest.filter(request_type=RequestType.FRIEND).count()
g_count = await FgRequest.filter(request_type=RequestType.GROUP).count()
data = {
"friend_count": f_count,
"group_count": g_count,
}
return Result.ok(data, f"{NICKNAME}带来了最新的数据!")
@router.get(
"/get_request_list", dependencies=[authentication()], description="获取请求列表"
)
async def _() -> Result:
try:
req_result = ReqResult()
data_list = await FgRequest.filter(handle_type__not_isnull=True).all()
for req in data_list:
if req.request_type == RequestType.FRIEND:
req_result.friend.append(
FriendRequestResult(
oid=req.id,
bot_id=req.bot_id,
id=req.user_id,
flag=req.flag,
nickname=req.nickname,
comment=req.comment,
ava_url=AVA_URL.format(req.user_id),
type=str(req.request_type).lower(),
)
)
else:
req_result.group.append(
GroupRequestResult(
oid=req.id,
bot_id=req.bot_id,
id=req.user_id,
flag=req.flag,
nickname=req.nickname,
comment=req.comment,
ava_url=GROUP_AVA_URL.format(req.group_id, req.group_id),
type=str(req.request_type).lower(),
invite_group=req.group_id,
group_name=None,
)
)
req_result.friend.reverse()
req_result.group.reverse()
except Exception as e:
logger.error("调用API错误", "/get_request", e=e)
return Result.fail(f"{type(e)}: {e}")
return Result.ok(req_result, f"{NICKNAME}带来了最新的数据!")
@router.delete(
"/clear_request", dependencies=[authentication()], description="清空请求列表"
)
async def _(request_type: Literal["private", "group"]) -> Result:
await FgRequest.filter(handle_type__not_isnull=True).update(
handle_type=RequestHandleType.IGNORE
)
return Result.ok(info="成功清除了数据!")
@router.post("/refuse_request", dependencies=[authentication()], description="拒绝请求")
async def _(parma: HandleRequest) -> Result:
try:
if bots := nonebot.get_bots():
bot_id = parma.bot_id
if bot_id not in nonebot.get_bots():
return Result.warning_("指定Bot未连接...")
try:
await FgRequest.refused(bots[bot_id], parma.id)
except ActionFailed as e:
await FgRequest.expire(parma.id)
return Result.warning_("请求失败,可能该请求已失效或请求数据错误...")
except NotFoundError:
return Result.warning_("未找到此Id请求...")
return Result.ok(info="成功处理了请求!")
return Result.warning_("无Bot连接...")
except Exception as e:
logger.error("调用API错误", "/refuse_request", e=e)
return Result.fail(f"{type(e)}: {e}")
@router.post("/delete_request", dependencies=[authentication()], description="忽略请求")
async def _(parma: HandleRequest) -> Result:
await FgRequest.expire(parma.id)
return Result.ok(info="成功处理了请求!")
@router.post(
"/approve_request", dependencies=[authentication()], description="同意请求"
)
async def _(parma: HandleRequest) -> Result:
try:
if bots := nonebot.get_bots():
bot_id = parma.bot_id
if bot_id not in nonebot.get_bots():
return Result.warning_("指定Bot未连接...")
if parma.request_type == "group":
if req := await FgRequest.get_or_none(id=parma.id):
if group := await GroupConsole.get_or_none(group_id=req.group_id):
await group.update_or_create(group_flag=1)
else:
group_info = await bots[bot_id].get_group_info(
group_id=req.group_id
)
await GroupConsole.update_or_create(
group_id=str(group_info["group_id"]),
defaults={
"group_name": group_info["group_name"],
"max_member_count": group_info["max_member_count"],
"member_count": group_info["member_count"],
"group_flag": 1,
},
)
else:
return Result.warning_("未找到此Id请求...")
try:
await FgRequest.approve(bots[bot_id], parma.id)
return Result.ok(info="成功处理了请求!")
except ActionFailed as e:
await FgRequest.expire(parma.id)
return Result.warning_("请求失败,可能该请求已失效或请求数据错误...")
return Result.warning_("无Bot连接...")
except Exception as e:
logger.error("调用API错误", "/approve_request", e=e)
return Result.fail(f"{type(e)}: {e}")
@router.post("/leave_group", dependencies=[authentication()], description="退群")
async def _(param: LeaveGroup) -> Result:
try:
if bots := nonebot.get_bots():
bot_id = param.bot_id
platform = PlatformUtils.get_platform(bots[bot_id])
if platform != "qq":
return Result.warning_("该平台不支持退群操作...")
group_list = await bots[bot_id].get_group_list()
if param.group_id not in [str(g["group_id"]) for g in group_list]:
return Result.warning_("Bot未在该群聊中...")
await bots[bot_id].set_group_leave(group_id=param.group_id)
return Result.ok(info="成功处理了请求!")
return Result.warning_("无Bot连接...")
except Exception as e:
logger.error("调用API错误", "/leave_group", e=e)
return Result.fail(f"{type(e)}: {e}")
@router.post("/delete_friend", dependencies=[authentication()], description="删除好友")
async def _(param: DeleteFriend) -> Result:
try:
if bots := nonebot.get_bots():
bot_id = param.bot_id
platform = PlatformUtils.get_platform(bots[bot_id])
if platform != "qq":
return Result.warning_("该平台不支持删除好友操作...")
friend_list = await bots[bot_id].get_friend_list()
if param.user_id not in [str(g["user_id"]) for g in friend_list]:
return Result.warning_("Bot未有其好友...")
await bots[bot_id].delete_friend(user_id=param.user_id)
return Result.ok(info="成功处理了请求!")
return Result.warning_("Bot未连接...")
except Exception as e:
logger.error("调用API错误", "/delete_friend", e=e)
return Result.fail(f"{type(e)}: {e}")
@router.get(
"/get_friend_detail", dependencies=[authentication()], description="获取好友详情"
)
async def _(bot_id: str, user_id: str) -> Result:
if bots := nonebot.get_bots():
if bot_id in bots:
if fd := [
x
for x in await bots[bot_id].get_friend_list()
if str(x["user_id"]) == user_id
]:
like_plugin_list = (
await Statistics.filter(user_id=user_id)
.annotate(count=Count("id"))
.group_by("plugin_name")
.order_by("-count")
.limit(5)
.values_list("plugin_name", "count")
)
like_plugin = {}
module_list = [x[0] for x in like_plugin_list]
plugins = await PluginInfo.filter(module__in=module_list).all()
module2name = {p.module: p.name for p in plugins}
for data in like_plugin_list:
name = module2name.get(data[0]) or data[0]
like_plugin[name] = data[1]
user = fd[0]
user_detail = UserDetail(
user_id=user_id,
ava_url=AVA_URL.format(user_id),
nickname=user["nickname"],
remark=user["remark"],
is_ban=await BanConsole.is_ban(user_id),
chat_count=await ChatHistory.filter(user_id=user_id).count(),
call_count=await Statistics.filter(user_id=user_id).count(),
like_plugin=like_plugin,
)
return Result.ok(user_detail)
else:
return Result.warning_("未添加指定好友...")
return Result.warning_("无Bot连接...")
@router.get(
"/get_group_detail", dependencies=[authentication()], description="获取群组详情"
)
async def _(bot_id: str, group_id: str) -> Result:
if bots := nonebot.get_bots():
if bot_id in bots:
group = await GroupConsole.get_or_none(group_id=group_id)
if not group:
return Result.warning_("指定群组未被收录...")
like_plugin_list = (
await Statistics.filter(group_id=group_id)
.annotate(count=Count("id"))
.group_by("plugin_name")
.order_by("-count")
.limit(5)
.values_list("plugin_name", "count")
)
like_plugin = {}
plugins = await PluginInfo.all()
module2name = {p.module: p.name for p in plugins}
for data in like_plugin_list:
name = module2name.get(data[0]) or data[0]
like_plugin[name] = data[1]
close_plugins = []
if group.block_plugin:
for module in group.block_plugin.split(","):
module_ = module.replace(":super", "")
is_super_block = module.endswith(":super")
plugin = Plugin(
module=module_,
plugin_name=module,
is_super_block=is_super_block,
)
plugin.plugin_name = module2name.get(module) or module
close_plugins.append(plugin)
all_task = await TaskInfo.annotate().values_list("module", "name")
task_module2name = {x[0]: x[1] for x in all_task}
task_list = []
if group.block_task:
split_task = group.block_task.split(",")
for task in all_task:
task_list.append(
Task(
name=task[0],
zh_name=task_module2name.get(task[0]) or task[0],
status=task[0] not in split_task,
)
)
group_detail = GroupDetail(
group_id=group_id,
ava_url=GROUP_AVA_URL.format(group_id, group_id),
name=group.group_name,
member_count=group.member_count,
max_member_count=group.max_member_count,
chat_count=await ChatHistory.filter(group_id=group_id).count(),
call_count=await Statistics.filter(group_id=group_id).count(),
like_plugin=like_plugin,
level=group.level,
status=group.status,
close_plugins=close_plugins,
task=task_list,
)
return Result.ok(group_detail)
else:
return Result.warning_("未添加指定群组...")
return Result.warning_("无Bot连接...")
@router.post(
"/send_message", dependencies=[authentication()], description="获取群组详情"
)
async def _(param: SendMessage) -> Result:
if bots := nonebot.get_bots():
if param.bot_id in bots:
platform = PlatformUtils.get_platform(bots[param.bot_id])
if platform != "qq":
return Result.warning_("暂不支持该平台...")
try:
if param.user_id:
await bots[param.bot_id].send_private_msg(
user_id=str(param.user_id), message=param.message
)
else:
await bots[param.bot_id].send_group_msg(
group_id=str(param.group_id), message=param.message
)
except Exception as e:
return Result.fail(str(e))
return Result.ok("发送成功!")
return Result.warning_("指定Bot未连接...")
return Result.warning_("无Bot连接...")
MSG_LIST = []
ID2NAME = {}
async def message_handle(
sub_log: str, type: Literal["private", "group"]
) -> Message | None:
global MSG_LIST, ID2NAME
pattern = PRIVATE_PATTERN if type == "private" else GROUP_PATTERN
msg_id = None
uid = None
gid = None
msg = None
img_list = re.findall(IMAGE_PATTERN, sub_log)
if r := re.search(pattern, sub_log):
if type == "private":
msg_id = r.group(1)
uid = r.group(2)
msg = r.group(3)
if uid not in ID2NAME:
if user := await FriendUser.get_or_none(user_id=uid):
ID2NAME[uid] = user.user_name or user.nickname
else:
msg_id = r.group(1)
uid = r.group(2)
gid = r.group(3)
msg = r.group(4)
if gid not in ID2NAME:
if user := await GroupInfoUser.get_or_none(user_id=uid, group_id=gid):
ID2NAME[uid] = user.user_name or user.nickname
if at_list := re.findall(AT_PATTERN, msg):
user_list = await GroupInfoUser.filter(
user_id__in=at_list, group_id=gid
).all()
id2name = {u.user_id: (u.user_name or u.nickname) for u in user_list}
for qq in at_list:
msg = re.sub(rf"\[CQ:at,qq={qq}\]", f"@{id2name[qq] or ''}", msg)
if msg_id in MSG_LIST:
return
MSG_LIST.append(msg_id)
messages = []
if msg and uid:
rep = re.split(r"\[CQ:image.*\]", msg)
if img_list:
for i in range(len(rep)):
messages.append(MessageItem(type="text", msg=rep[i]))
if i < len(img_list):
messages.append(MessageItem(type="img", msg=img_list[i]))
else:
messages = [MessageItem(type="text", msg=x) for x in rep]
return Message(
object_id=uid if type == "private" else gid, # type: ignore
user_id=uid,
group_id=gid,
message=messages,
name=ID2NAME.get(uid) or "",
ava_url=AVA_URL.format(uid),
)
return None
@ws_router.websocket("/chat")
async def _(websocket: WebSocket):
await websocket.accept()
async def log_listener(log: str):
global MSG_LIST, ID2NAME
sub_log = re.sub(SUB_PATTERN, "", log)
img_list = re.findall(IMAGE_PATTERN, sub_log)
if "message.private.friend" in log:
if message := await message_handle(sub_log, "private"):
await websocket.send_json(message.dict())
else:
if r := re.search(GROUP_PATTERN, sub_log):
if message := await message_handle(sub_log, "group"):
await websocket.send_json(message.dict())
if len(MSG_LIST) > 30:
MSG_LIST = MSG_LIST[-1:]
LOG_STORAGE.listeners.add(log_listener)
try:
while websocket.client_state == WebSocketState.CONNECTED:
recv = await websocket.receive()
except WebSocketDisconnect:
pass
finally:
LOG_STORAGE.listeners.remove(log_listener)
return
@@ -0,0 +1,265 @@
from typing import Literal
from pydantic import BaseModel
class Group(BaseModel):
"""
群组信息
"""
group_id: str
"""群组id"""
group_name: str
"""群组名称"""
member_count: int
"""成员人数"""
max_member_count: int
"""群组最大人数"""
class Task(BaseModel):
"""
被动技能
"""
name: str
"""被动名称"""
zh_name: str
"""被动中文名称"""
status: bool
"""状态"""
class Plugin(BaseModel):
"""
插件
"""
module: str
"""模块名"""
plugin_name: str
"""中文名"""
is_super_block: bool
"""是否超级用户禁用"""
class GroupResult(BaseModel):
"""
群组返回数据
"""
group_id: str
"""群组id"""
group_name: str
"""群组名称"""
ava_url: str
"""群组头像"""
class Friend(BaseModel):
"""
好友数据
"""
user_id: str
"""用户id"""
nickname: str = ""
"""昵称"""
remark: str = ""
"""备注"""
ava_url: str = ""
"""头像url"""
class UpdateGroup(BaseModel):
"""
更新群组信息
"""
group_id: str
"""群号"""
status: bool
"""状态"""
level: int
"""群权限"""
task: list[str]
"""被动状态"""
close_plugins: list[str]
"""关闭插件"""
class FriendRequestResult(BaseModel):
"""
好友/群组请求管理
"""
bot_id: str
"""bot_id"""
oid: int
"""排序"""
id: str
"""id"""
flag: str
"""flag"""
nickname: str | None
"""昵称"""
comment: str | None
"""备注信息"""
ava_url: str
"""头像"""
type: str
"""类型 private group"""
class GroupRequestResult(FriendRequestResult):
"""
群聊邀请请求
"""
invite_group: str
"""邀请群聊"""
group_name: str | None
"""群聊名称"""
class HandleRequest(BaseModel):
"""
操作请求接收数据
"""
bot_id: str | None = None
"""bot_id"""
id: int
"""数据id"""
request_type: Literal["private", "group"]
"""类型"""
class LeaveGroup(BaseModel):
"""
退出群聊
"""
bot_id: str
"""bot_id"""
group_id: str
"""群聊id"""
class DeleteFriend(BaseModel):
"""
删除好友
"""
bot_id: str
"""bot_id"""
user_id: str
"""用户id"""
class ReqResult(BaseModel):
"""
好友/群组请求列表
"""
friend: list[FriendRequestResult] = []
"""好友请求列表"""
group: list[GroupRequestResult] = []
"""群组请求列表"""
class UserDetail(BaseModel):
"""
用户详情
"""
user_id: str
"""用户id"""
ava_url: str
"""头像url"""
nickname: str
"""昵称"""
remark: str
"""备注"""
is_ban: bool
"""是否被ban"""
chat_count: int
"""发言次数"""
call_count: int
"""功能调用次数"""
like_plugin: dict[str, int]
"""最喜爱的功能"""
class GroupDetail(BaseModel):
"""
用户详情
"""
group_id: str
"""群组id"""
ava_url: str
"""头像url"""
name: str
"""名称"""
member_count: int
"""成员数"""
max_member_count: int
"""最大成员数"""
chat_count: int
"""发言次数"""
call_count: int
"""功能调用次数"""
like_plugin: dict[str, int]
"""最喜爱的功能"""
level: int
"""群权限"""
status: bool
"""状态(睡眠)"""
close_plugins: list[Plugin]
"""关闭的插件"""
task: list[Task]
"""被动列表"""
class MessageItem(BaseModel):
type: str
"""消息类型"""
msg: str
"""内容"""
class Message(BaseModel):
"""
消息
"""
object_id: str
"""主体id user_id 或 group_id"""
user_id: str
"""用户id"""
group_id: str | None = None
"""群组id"""
message: list[MessageItem]
"""消息"""
name: str
"""用户名称"""
ava_url: str
"""用户头像"""
class SendMessage(BaseModel):
"""
发送消息
"""
bot_id: str
"""bot id"""
user_id: str | None = None
"""用户id"""
group_id: str | None = None
"""群组id"""
message: str
"""消息"""
@@ -0,0 +1,187 @@
import re
import cattrs
from fastapi import APIRouter, Query
from zhenxun.configs.config import Config
from zhenxun.models.plugin_info import PluginInfo as DbPluginInfo
from zhenxun.services.log import logger
from zhenxun.utils.enum import BlockType, PluginType
from ....base_model import Result
from ....utils import authentication
from .model import (
PluginConfig,
PluginCount,
PluginDetail,
PluginInfo,
PluginSwitch,
UpdatePlugin,
)
router = APIRouter(prefix="/plugin")
@router.get(
"/get_plugin_list", dependencies=[authentication()], deprecated="获取插件列表" # type: ignore
)
async def _(
plugin_type: list[PluginType] = Query(None), menu_type: str | None = None
) -> Result:
try:
plugin_list: list[PluginInfo] = []
query = DbPluginInfo
if plugin_type:
query = query.filter(plugin_type__in=plugin_type)
if menu_type:
query = query.filter(menu_type=menu_type)
plugins = await query.all()
for plugin in plugins:
plugin_info = PluginInfo(
module=plugin.module,
plugin_name=plugin.name,
default_status=plugin.default_status,
limit_superuser=plugin.limit_superuser,
cost_gold=plugin.cost_gold,
menu_type=plugin.menu_type,
version=plugin.version or 0,
level=plugin.level,
status=plugin.status,
author=plugin.author,
)
plugin_list.append(plugin_info)
except Exception as e:
logger.error("调用API错误", "/get_plugins", e=e)
return Result.fail(f"{type(e)}: {e}")
return Result.ok(plugin_list, "拿到了新鲜出炉的数据!")
@router.get(
"/get_plugin_count", dependencies=[authentication()], deprecated="获取插件数量" # type: ignore
)
async def _() -> Result:
plugin_count = PluginCount()
plugin_count.normal = await DbPluginInfo.filter(
plugin_type=PluginType.NORMAL
).count()
plugin_count.admin = await DbPluginInfo.filter(
plugin_type__in=[PluginType.ADMIN, PluginType.SUPER_AND_ADMIN]
).count()
plugin_count.superuser = await DbPluginInfo.filter(
plugin_type=[PluginType.SUPERUSER, PluginType.SUPER_AND_ADMIN]
).count()
plugin_count.other = await DbPluginInfo.filter(
plugin_type=PluginType.HIDDEN
).count()
return Result.ok(plugin_count)
@router.post(
"/update_plugin", dependencies=[authentication()], description="更新插件参数"
)
async def _(plugin: UpdatePlugin) -> Result:
try:
db_plugin = await DbPluginInfo.get_or_none(module=plugin.module)
if not db_plugin:
return Result.fail("插件不存在...")
db_plugin.default_status = plugin.default_status
db_plugin.limit_superuser = plugin.limit_superuser
db_plugin.cost_gold = plugin.cost_gold
db_plugin.level = plugin.level
db_plugin.menu_type = plugin.menu_type
db_plugin.block_type = plugin.block_type
if plugin.block_type == BlockType.ALL:
db_plugin.status = False
else:
db_plugin.status = True
await db_plugin.save()
# 配置项
if plugin.configs and (configs := Config.get(plugin.module)):
for key in plugin.configs:
if c := configs.configs.get(key):
value = plugin.configs[key]
if c.type and value is not None:
value = cattrs.structure(value, c.type)
Config.set_config(plugin.module, key, value)
except Exception as e:
logger.error("调用API错误", "/update_plugins", e=e)
return Result.fail(f"{type(e)}: {e}")
return Result.ok(info="已经帮你写好啦!")
@router.post("/change_switch", dependencies=[authentication()], description="开关插件")
async def _(param: PluginSwitch) -> Result:
db_plugin = await DbPluginInfo.get_or_none(module=param.module)
if not db_plugin:
return Result.fail("插件不存在...")
if not param.status:
db_plugin.block_type = BlockType.ALL
db_plugin.status = False
else:
db_plugin.block_type = None
db_plugin.status = True
await db_plugin.save()
return Result.ok(info="成功改变了开关状态!")
@router.get(
"/get_plugin_menu_type", dependencies=[authentication()], description="获取插件类型"
)
async def _() -> Result:
menu_type_list = []
result = await DbPluginInfo.annotate().values_list("menu_type", flat=True)
for r in result:
if r not in menu_type_list:
menu_type_list.append(r)
return Result.ok(menu_type_list)
@router.get("/get_plugin", dependencies=[authentication()], description="获取插件详情")
async def _(module: str) -> Result:
db_plugin = await DbPluginInfo.get_or_none(module=module)
if not db_plugin:
return Result.fail("插件不存在...")
config_list = []
if config := Config.get(module):
for cfg in config.configs:
type_str = ""
type_inner = None
x = str(config.configs[cfg].type)
r = re.search(r"<class '(.*)'>", str(config.configs[cfg].type))
if r:
type_str = r.group(1)
else:
r = re.search(r"typing\.(.*)\[(.*)\]", str(config.configs[cfg].type))
if r:
type_str = r.group(1)
if type_str:
type_str = type_str.lower()
type_inner = r.group(2)
if type_inner:
type_inner = [x.strip() for x in type_inner.split(",")]
config_list.append(
PluginConfig(
module=module,
key=cfg,
value=config.configs[cfg].value,
help=config.configs[cfg].help,
default_value=config.configs[cfg].default_value,
type=type_str,
type_inner=type_inner, # type: ignore
)
)
plugin_info = PluginDetail(
module=module,
plugin_name=db_plugin.name,
default_status=db_plugin.default_status,
limit_superuser=db_plugin.limit_superuser,
cost_gold=db_plugin.cost_gold,
menu_type=db_plugin.menu_type,
version=db_plugin.version or "0",
level=db_plugin.level,
status=db_plugin.status,
author=db_plugin.author,
config_list=config_list,
block_type=db_plugin.block_type,
)
return Result.ok(plugin_info)
@@ -0,0 +1,125 @@
from typing import Any
from pydantic import BaseModel
from zhenxun.utils.enum import BlockType
class PluginSwitch(BaseModel):
"""
插件开关
"""
module: str
"""模块"""
status: bool
"""开关状态"""
class UpdateConfig(BaseModel):
"""
配置项修改参数
"""
module: str
"""模块"""
key: str
"""配置项key"""
value: Any
"""配置项值"""
class UpdatePlugin(BaseModel):
"""
插件修改参数
"""
module: str
"""模块"""
default_status: bool
"""默认开关"""
limit_superuser: bool
"""限制超级用户"""
cost_gold: int
"""金币花费"""
menu_type: str
"""插件菜单类型"""
level: int
"""插件所需群权限"""
block_type: BlockType | None = None
"""禁用类型"""
configs: dict[str, Any] | None = None
"""配置项"""
class PluginInfo(BaseModel):
"""
基本插件信息
"""
module: str
"""插件名称"""
plugin_name: str
"""插件中文名称"""
default_status: bool
"""默认开关"""
limit_superuser: bool
"""限制超级用户"""
cost_gold: int
"""花费金币"""
menu_type: str
"""插件菜单类型"""
version: str
"""插件版本"""
level: int
"""群权限"""
status: bool
"""当前状态"""
author: str | None = None
"""作者"""
block_type: BlockType | None = None
"""禁用类型"""
class PluginConfig(BaseModel):
"""
插件配置项
"""
module: str
"""模块"""
key: str
"""键"""
value: Any
"""值"""
help: str | None = None
"""帮助"""
default_value: Any
"""默认值"""
type: Any = None
"""值类型"""
type_inner: list[str] | None = None
"""List Tuple等内部类型检验"""
class PluginCount(BaseModel):
"""
插件数量
"""
normal: int = 0
"""普通插件"""
admin: int = 0
"""管理员插件"""
superuser: int = 0
"""超级用户插件"""
other: int = 0
"""其他插件"""
class PluginDetail(PluginInfo):
"""
插件详情
"""
config_list: list[PluginConfig]
@@ -0,0 +1,121 @@
import os
import shutil
from pathlib import Path
from typing import List, Optional
from fastapi import APIRouter
from ....base_model import Result
from ....utils import authentication, get_system_disk
from .model import AddFile, DeleteFile, DirFile, RenameFile, SaveFile
router = APIRouter(prefix="/system")
@router.get("/get_dir_list", dependencies=[authentication()], description="获取文件列表")
async def _(path: Optional[str] = None) -> Result:
base_path = Path(path) if path else Path()
data_list = []
for file in os.listdir(base_path):
data_list.append(DirFile(is_file=not (base_path / file).is_dir(), name=file, parent=path))
return Result.ok(data_list)
@router.get("/get_resources_size", dependencies=[authentication()], description="获取文件列表")
async def _(full_path: Optional[str] = None) -> Result:
return Result.ok(await get_system_disk(full_path))
@router.post("/delete_file", dependencies=[authentication()], description="删除文件")
async def _(param: DeleteFile) -> Result:
path = Path(param.full_path)
if not path or not path.exists():
return Result.warning_("文件不存在...")
try:
path.unlink()
return Result.ok('删除成功!')
except Exception as e:
return Result.warning_('删除失败: ' + str(e))
@router.post("/delete_folder", dependencies=[authentication()], description="删除文件夹")
async def _(param: DeleteFile) -> Result:
path = Path(param.full_path)
if not path or not path.exists() or path.is_file():
return Result.warning_("文件夹不存在...")
try:
shutil.rmtree(path.absolute())
return Result.ok('删除成功!')
except Exception as e:
return Result.warning_('删除失败: ' + str(e))
@router.post("/rename_file", dependencies=[authentication()], description="重命名文件")
async def _(param: RenameFile) -> Result:
path = (Path(param.parent) / param.old_name) if param.parent else Path(param.old_name)
if not path or not path.exists():
return Result.warning_("文件不存在...")
try:
path.rename(path.parent / param.name)
return Result.ok('重命名成功!')
except Exception as e:
return Result.warning_('重命名失败: ' + str(e))
@router.post("/rename_folder", dependencies=[authentication()], description="重命名文件夹")
async def _(param: RenameFile) -> Result:
path = (Path(param.parent) / param.old_name) if param.parent else Path(param.old_name)
if not path or not path.exists() or path.is_file():
return Result.warning_("文件夹不存在...")
try:
new_path = path.parent / param.name
shutil.move(path.absolute(), new_path.absolute())
return Result.ok('重命名成功!')
except Exception as e:
return Result.warning_('重命名失败: ' + str(e))
@router.post("/add_file", dependencies=[authentication()], description="新建文件")
async def _(param: AddFile) -> Result:
path = (Path(param.parent) / param.name) if param.parent else Path(param.name)
if path.exists():
return Result.warning_("文件已存在...")
try:
path.open('w')
return Result.ok('新建文件成功!')
except Exception as e:
return Result.warning_('新建文件失败: ' + str(e))
@router.post("/add_folder", dependencies=[authentication()], description="新建文件夹")
async def _(param: AddFile) -> Result:
path = (Path(param.parent) / param.name) if param.parent else Path(param.name)
if path.exists():
return Result.warning_("文件夹已存在...")
try:
path.mkdir()
return Result.ok('新建文件夹成功!')
except Exception as e:
return Result.warning_('新建文件夹失败: ' + str(e))
@router.get("/read_file", dependencies=[authentication()], description="读取文件")
async def _(full_path: str) -> Result:
path = Path(full_path)
if not path.exists():
return Result.warning_("文件不存在...")
try:
text = path.read_text(encoding='utf-8')
return Result.ok(text)
except Exception as e:
return Result.warning_('新建文件夹失败: ' + str(e))
@router.post("/save_file", dependencies=[authentication()], description="读取文件")
async def _(param: SaveFile) -> Result:
path = Path(param.full_path)
try:
with path.open('w') as f:
f.write(param.content)
return Result.ok("更新成功!")
except Exception as e:
return Result.warning_('新建文件夹失败: ' + str(e))
@@ -0,0 +1,64 @@
from datetime import datetime
from typing import Literal, Optional
from pydantic import BaseModel
class DirFile(BaseModel):
"""
文件或文件夹
"""
is_file: bool
"""是否为文件"""
name: str
"""文件夹或文件名称"""
parent: Optional[str] = None
"""父级"""
class DeleteFile(BaseModel):
"""
删除文件
"""
full_path: str
"""文件全路径"""
class RenameFile(BaseModel):
"""
删除文件
"""
parent: Optional[str]
"""父路径"""
old_name: str
"""旧名称"""
name: str
"""新名称"""
class AddFile(BaseModel):
"""
新建文件
"""
parent: Optional[str]
"""父路径"""
name: str
"""新名称"""
class SaveFile(BaseModel):
"""
保存文件
"""
full_path: str
"""全路径"""
content: str
"""内容"""