369 lines
17 KiB
Python
369 lines
17 KiB
Python
import datetime
|
|||
|
|
from typing import Optional, Union
|
||
|
|
|
||
|
|
from fastapi import FastAPI
|
||
|
|
from fastapi import Header, HTTPException, Depends
|
||
|
|
from fastapi.responses import JSONResponse, HTMLResponse, RedirectResponse
|
||
|
|
from jose import jwt
|
||
|
|
from nonebot import get_app, get_adapter, logger
|
||
|
|
from nonebot.adapters.onebot.v11 import Adapter
|
||
|
|
from pydantic import BaseModel
|
||
|
|
from sqlalchemy import select, delete, update, func
|
||
|
|
|
||
|
|
try:
|
||
|
|
import jieba_fast as jieba
|
||
|
|
except ImportError:
|
||
|
|
import jieba
|
||
|
|
|
||
|
|
from nonebot_plugin_orm import get_session
|
||
|
|
from .handler import LearningChat
|
||
|
|
from .models import ChatMessage, ChatContext, ChatAnswer, ChatBlackList
|
||
|
|
from .config import config_manager, driver
|
||
|
|
from .web_frontend import ensure_frontend_ready, mount_frontend
|
||
|
|
|
||
|
|
|
||
|
|
def authentication():
|
||
|
|
def inner(token: Optional[str] = Header(...)):
|
||
|
|
try:
|
||
|
|
payload = jwt.decode(token, config_manager.config.web_secret_key, algorithms="HS256")
|
||
|
|
if (
|
||
|
|
not (username := payload.get("username"))
|
||
|
|
or username != config_manager.config.web_username
|
||
|
|
):
|
||
|
|
raise HTTPException(status_code=400, detail="登录验证失败或已失效,请重新登录")
|
||
|
|
except (jwt.JWTError, jwt.ExpiredSignatureError, AttributeError):
|
||
|
|
raise HTTPException(status_code=400, detail="登录验证失败或已失效,请重新登录")
|
||
|
|
|
||
|
|
return Depends(inner)
|
||
|
|
|
||
|
|
|
||
|
|
class UserModel(BaseModel):
|
||
|
|
username: str
|
||
|
|
password: str
|
||
|
|
|
||
|
|
|
||
|
|
@driver.on_startup
|
||
|
|
async def init_web():
|
||
|
|
if not config_manager.config.enable_web:
|
||
|
|
return
|
||
|
|
if not await ensure_frontend_ready():
|
||
|
|
logger.warning("群聊学习 | 前端构建失败, 管理页面将不可用")
|
||
|
|
app: FastAPI = get_app()
|
||
|
|
|
||
|
|
@app.post("/learning_chat/api/login", response_class=JSONResponse)
|
||
|
|
async def login(user: UserModel):
|
||
|
|
if (
|
||
|
|
user.username != config_manager.config.web_username
|
||
|
|
or user.password != config_manager.config.web_password
|
||
|
|
):
|
||
|
|
return {"status": -100, "msg": "登录失败,请确认用户ID和密码无误"}
|
||
|
|
token = jwt.encode(
|
||
|
|
{
|
||
|
|
"username": user.username,
|
||
|
|
"exp": datetime.datetime.now(datetime.timezone.utc) + datetime.timedelta(minutes=30),
|
||
|
|
},
|
||
|
|
config_manager.config.web_secret_key,
|
||
|
|
algorithm="HS256",
|
||
|
|
)
|
||
|
|
return {"status": 0, "msg": "登录成功", "data": {"token": token}}
|
||
|
|
|
||
|
|
@app.options("/learning_chat/api/login")
|
||
|
|
async def options_login():
|
||
|
|
return JSONResponse(content={}, status_code=200)
|
||
|
|
|
||
|
|
@app.get("/learning_chat/api/get_group_list", response_class=JSONResponse, dependencies=[authentication()])
|
||
|
|
async def get_group_list_api():
|
||
|
|
try:
|
||
|
|
bots = get_adapter(Adapter).bots
|
||
|
|
if len(bots) == 0:
|
||
|
|
return {"status": -100, "msg": "获取群和好友列表失败,请确认已连接GOCQ"}
|
||
|
|
bot = list(bots.values())[0]
|
||
|
|
group_list = await bot.get_group_list()
|
||
|
|
group_list = [
|
||
|
|
{"label": f'{group["group_name"]}({group["group_id"]})', "value": group["group_id"]}
|
||
|
|
for group in group_list
|
||
|
|
]
|
||
|
|
return {"status": 0, "msg": "ok", "data": {"group_list": group_list}}
|
||
|
|
except ValueError:
|
||
|
|
return {"status": -100, "msg": "获取群和好友列表失败,请确认已连接GOCQ"}
|
||
|
|
|
||
|
|
@app.options("/learning_chat/api/get_group_list")
|
||
|
|
async def options_get_group_list():
|
||
|
|
return JSONResponse(content={}, status_code=200)
|
||
|
|
|
||
|
|
@app.get("/learning_chat/api/chat_global_config", response_class=JSONResponse, dependencies=[authentication()])
|
||
|
|
async def get_chat_global_config():
|
||
|
|
# 注意: 不再注入 member_list —— 前端全局配置页不使用(仅分群配置需要,
|
||
|
|
# 由 chat_group_config 单独拉取单个群成员), 避免串行拉取全部群成员导致接口卡顿
|
||
|
|
return config_manager.config.dict(exclude={"group_config"})
|
||
|
|
|
||
|
|
@app.post("/learning_chat/api/chat_global_config", response_class=JSONResponse, dependencies=[authentication()])
|
||
|
|
async def post_chat_global_config(data: dict):
|
||
|
|
config_manager.config.update(**data)
|
||
|
|
config_manager.save()
|
||
|
|
async with get_session() as session:
|
||
|
|
await session.execute(
|
||
|
|
update(ChatContext)
|
||
|
|
.where(ChatContext.count > config_manager.config.learn_max_count)
|
||
|
|
.values(count=config_manager.config.learn_max_count)
|
||
|
|
)
|
||
|
|
await session.execute(
|
||
|
|
update(ChatAnswer)
|
||
|
|
.where(ChatAnswer.count > config_manager.config.learn_max_count)
|
||
|
|
.values(count=config_manager.config.learn_max_count)
|
||
|
|
)
|
||
|
|
await session.commit()
|
||
|
|
jieba.load_userdict(config_manager.config.dictionary)
|
||
|
|
return {"status": 0, "msg": "保存成功"}
|
||
|
|
|
||
|
|
@app.options("/learning_chat/api/chat_global_config")
|
||
|
|
async def options_chat_global_config():
|
||
|
|
return JSONResponse(content={}, status_code=200)
|
||
|
|
|
||
|
|
@app.get("/learning_chat/api/chat_group_config", response_class=JSONResponse, dependencies=[authentication()])
|
||
|
|
async def get_chat_group_config(group_id: int):
|
||
|
|
try:
|
||
|
|
bots = get_adapter(Adapter).bots
|
||
|
|
if len(bots) == 0:
|
||
|
|
return {"status": -100, "msg": "获取群和好友列表失败,请确认已连接GOCQ"}
|
||
|
|
bot = list(bots.values())[0]
|
||
|
|
members = await bot.get_group_member_list(group_id=group_id)
|
||
|
|
member_list = [
|
||
|
|
{"label": f'{member["nickname"] or member["card"]}({member["user_id"]})', "value": member["user_id"]}
|
||
|
|
for member in members
|
||
|
|
]
|
||
|
|
config = config_manager.get_group_config(group_id).dict()
|
||
|
|
config["break_probability"] = config["break_probability"] * 100
|
||
|
|
config["speak_continuously_probability"] = config["speak_continuously_probability"] * 100
|
||
|
|
config["speak_poke_probability"] = config["speak_poke_probability"] * 100
|
||
|
|
config["member_list"] = member_list
|
||
|
|
return config
|
||
|
|
except ValueError:
|
||
|
|
return {"status": -100, "msg": "获取群和好友列表失败,请确认已连接GOCQ"}
|
||
|
|
|
||
|
|
@app.post("/learning_chat/api/chat_group_config", response_class=JSONResponse, dependencies=[authentication()])
|
||
|
|
async def post_chat_group_config(group_id: Union[int, str], data: dict):
|
||
|
|
if not data.get("answer_threshold_weights"):
|
||
|
|
return {"status": 400, "msg": "回复阈值权重不能为空,必须至少有一个数值"}
|
||
|
|
data["break_probability"] = data["break_probability"] / 100
|
||
|
|
data["speak_continuously_probability"] = data["speak_continuously_probability"] / 100
|
||
|
|
data["speak_poke_probability"] = data["speak_poke_probability"] / 100
|
||
|
|
bots = get_adapter(Adapter).bots
|
||
|
|
if len(bots) == 0:
|
||
|
|
return {"status": -100, "msg": "获取群和好友列表失败,请确认已连接GOCQ"}
|
||
|
|
bot = list(bots.values())[0]
|
||
|
|
groups = (
|
||
|
|
[{"group_id": group_id}] if group_id != "all" else await bot.get_group_list()
|
||
|
|
)
|
||
|
|
for group in groups:
|
||
|
|
config = config_manager.get_group_config(int(group["group_id"]))
|
||
|
|
config.update(**data)
|
||
|
|
config_manager.config.group_config[int(group["group_id"])] = config
|
||
|
|
config_manager.save()
|
||
|
|
return {"status": 0, "msg": "保存成功"}
|
||
|
|
|
||
|
|
@app.options("/learning_chat/api/chat_group_config")
|
||
|
|
async def options_chat_group_config():
|
||
|
|
return JSONResponse(content={}, status_code=200)
|
||
|
|
|
||
|
|
@app.get("/learning_chat/api/get_chat_messages", response_class=JSONResponse, dependencies=[authentication()])
|
||
|
|
async def get_chat_messages(
|
||
|
|
page: int = 1,
|
||
|
|
perPage: int = 10,
|
||
|
|
orderBy: str = "time",
|
||
|
|
orderDir: str = "desc",
|
||
|
|
group_id: Optional[str] = None,
|
||
|
|
user_id: Optional[str] = None,
|
||
|
|
message: Optional[str] = None,
|
||
|
|
):
|
||
|
|
async with get_session() as session:
|
||
|
|
stmt = select(ChatMessage)
|
||
|
|
if group_id:
|
||
|
|
stmt = stmt.where(ChatMessage.group_id == int(group_id))
|
||
|
|
if user_id:
|
||
|
|
stmt = stmt.where(ChatMessage.user_id == int(user_id))
|
||
|
|
if message:
|
||
|
|
stmt = stmt.where(ChatMessage.raw_message.contains(message))
|
||
|
|
|
||
|
|
order_col = getattr(ChatMessage, orderBy or "time", ChatMessage.time)
|
||
|
|
stmt = stmt.order_by(order_col.asc() if (orderDir or "desc") == "asc" else order_col.desc())
|
||
|
|
|
||
|
|
count_result = await session.execute(select(func.count()).select_from(stmt.subquery()))
|
||
|
|
total = count_result.scalar()
|
||
|
|
|
||
|
|
items_result = await session.execute(stmt.offset((page - 1) * perPage).limit(perPage))
|
||
|
|
items = [
|
||
|
|
{c.name: getattr(row, c.name) for c in ChatMessage.__table__.columns}
|
||
|
|
for row in items_result.scalars().all()
|
||
|
|
]
|
||
|
|
return {"status": 0, "msg": "ok", "data": {"items": items, "total": total}}
|
||
|
|
|
||
|
|
@app.options("/learning_chat/api/get_chat_messages")
|
||
|
|
async def options_get_chat_messages():
|
||
|
|
return JSONResponse(content={}, status_code=200)
|
||
|
|
|
||
|
|
@app.get("/learning_chat/api/get_chat_contexts", response_class=JSONResponse, dependencies=[authentication()])
|
||
|
|
async def get_chat_context(
|
||
|
|
page: int = 1,
|
||
|
|
perPage: int = 10,
|
||
|
|
orderBy: str = "time",
|
||
|
|
orderDir: str = "desc",
|
||
|
|
keywords: Optional[str] = None,
|
||
|
|
):
|
||
|
|
async with get_session() as session:
|
||
|
|
stmt = select(ChatContext)
|
||
|
|
if keywords:
|
||
|
|
stmt = stmt.where(ChatContext.keywords.contains(keywords))
|
||
|
|
order_col = getattr(ChatContext, orderBy or "time", ChatContext.time)
|
||
|
|
stmt = stmt.order_by(order_col.asc() if (orderDir or "desc") == "asc" else order_col.desc())
|
||
|
|
|
||
|
|
count_result = await session.execute(select(func.count()).select_from(stmt.subquery()))
|
||
|
|
total = count_result.scalar()
|
||
|
|
|
||
|
|
items_result = await session.execute(stmt.offset((page - 1) * perPage).limit(perPage))
|
||
|
|
items = [
|
||
|
|
{c.name: getattr(row, c.name) for c in ChatContext.__table__.columns}
|
||
|
|
for row in items_result.scalars().all()
|
||
|
|
]
|
||
|
|
return {"status": 0, "msg": "ok", "data": {"items": items, "total": total}}
|
||
|
|
|
||
|
|
@app.options("/learning_chat/api/get_chat_contexts")
|
||
|
|
async def options_get_chat_contexts():
|
||
|
|
return JSONResponse(content={}, status_code=200)
|
||
|
|
|
||
|
|
@app.get("/learning_chat/api/get_chat_answers", response_class=JSONResponse, dependencies=[authentication()])
|
||
|
|
async def get_chat_answers(
|
||
|
|
context_id: Optional[int] = None,
|
||
|
|
page: int = 1,
|
||
|
|
perPage: int = 10,
|
||
|
|
orderBy: str = "count",
|
||
|
|
orderDir: str = "desc",
|
||
|
|
keywords: Optional[str] = None,
|
||
|
|
):
|
||
|
|
async with get_session() as session:
|
||
|
|
stmt = select(ChatAnswer)
|
||
|
|
if context_id:
|
||
|
|
stmt = stmt.where(ChatAnswer.context_id == context_id)
|
||
|
|
if keywords:
|
||
|
|
stmt = stmt.where(ChatAnswer.keywords.contains(keywords))
|
||
|
|
order_col = getattr(ChatAnswer, orderBy or "count", ChatAnswer.count)
|
||
|
|
stmt = stmt.order_by(order_col.asc() if (orderDir or "desc") == "asc" else order_col.desc())
|
||
|
|
|
||
|
|
count_result = await session.execute(select(func.count()).select_from(stmt.subquery()))
|
||
|
|
total = count_result.scalar()
|
||
|
|
|
||
|
|
items_result = await session.execute(stmt.offset((page - 1) * perPage).limit(perPage))
|
||
|
|
items = []
|
||
|
|
for row in items_result.scalars().all():
|
||
|
|
item = {c.name: getattr(row, c.name) for c in ChatAnswer.__table__.columns}
|
||
|
|
item["messages"] = [{"msg": m} for m in item["messages"]]
|
||
|
|
items.append(item)
|
||
|
|
return {"status": 0, "msg": "ok", "data": {"items": items, "total": total}}
|
||
|
|
|
||
|
|
@app.options("/learning_chat/api/get_chat_answers")
|
||
|
|
async def options_get_chat_answers():
|
||
|
|
return JSONResponse(content={}, status_code=200)
|
||
|
|
|
||
|
|
@app.get("/learning_chat/api/get_chat_blacklist", response_class=JSONResponse, dependencies=[authentication()])
|
||
|
|
async def get_chat_blacklist(
|
||
|
|
page: int = 1,
|
||
|
|
perPage: int = 10,
|
||
|
|
keywords: Optional[str] = None,
|
||
|
|
bans: Optional[str] = None,
|
||
|
|
):
|
||
|
|
async with get_session() as session:
|
||
|
|
stmt = select(ChatBlackList).order_by(ChatBlackList.id.desc())
|
||
|
|
if keywords:
|
||
|
|
stmt = stmt.where(ChatBlackList.keywords.contains(keywords))
|
||
|
|
|
||
|
|
count_result = await session.execute(select(func.count()).select_from(stmt.subquery()))
|
||
|
|
total = count_result.scalar()
|
||
|
|
|
||
|
|
# bans 过滤需要基于全部数据(而非当前页), 故先全量取回再过滤、再分页
|
||
|
|
items_result = await session.execute(stmt)
|
||
|
|
items = []
|
||
|
|
for row in items_result.scalars().all():
|
||
|
|
item = {c.name: getattr(row, c.name) for c in ChatBlackList.__table__.columns}
|
||
|
|
ban_ids = item["ban_group_id"] or []
|
||
|
|
item["bans"] = "全局禁用" if item["global_ban"] else (str(ban_ids[0]) if ban_ids else "")
|
||
|
|
items.append(item)
|
||
|
|
if bans:
|
||
|
|
items = [x for x in items if bans in x["bans"]]
|
||
|
|
total = len(items)
|
||
|
|
items = items[(page - 1) * perPage : page * perPage]
|
||
|
|
return {"status": 0, "msg": "ok", "data": {"items": items, "total": total}}
|
||
|
|
|
||
|
|
@app.options("/learning_chat/api/get_chat_blacklist")
|
||
|
|
async def options_get_chat_blacklist():
|
||
|
|
return JSONResponse(content={}, status_code=200)
|
||
|
|
|
||
|
|
@app.delete("/learning_chat/api/delete_chat", response_class=JSONResponse, dependencies=[authentication()])
|
||
|
|
async def delete_chat(id: int, type: str):
|
||
|
|
try:
|
||
|
|
async with get_session() as session:
|
||
|
|
if type == "message":
|
||
|
|
await session.execute(delete(ChatMessage).where(ChatMessage.id == id))
|
||
|
|
elif type == "context":
|
||
|
|
await session.execute(delete(ChatAnswer).where(ChatAnswer.context_id == id))
|
||
|
|
await session.execute(delete(ChatContext).where(ChatContext.id == id))
|
||
|
|
elif type == "answer":
|
||
|
|
await session.execute(delete(ChatAnswer).where(ChatAnswer.id == id))
|
||
|
|
elif type == "blacklist":
|
||
|
|
await session.execute(delete(ChatBlackList).where(ChatBlackList.id == id))
|
||
|
|
await session.commit()
|
||
|
|
return {"status": 0, "msg": "删除成功"}
|
||
|
|
except Exception as e:
|
||
|
|
return {"status": 500, "msg": f"删除失败,{e}"}
|
||
|
|
|
||
|
|
@app.options("/learning_chat/api/delete_chat")
|
||
|
|
async def options_delete_chat():
|
||
|
|
return JSONResponse(content={}, status_code=200)
|
||
|
|
|
||
|
|
@app.put("/learning_chat/api/ban_chat", response_class=JSONResponse, dependencies=[authentication()])
|
||
|
|
async def ban_chat(id: int, type: str):
|
||
|
|
try:
|
||
|
|
async with get_session() as session:
|
||
|
|
if type == "message":
|
||
|
|
result = await session.execute(select(ChatMessage).where(ChatMessage.id == id))
|
||
|
|
data = result.scalar_one()
|
||
|
|
elif type == "context":
|
||
|
|
result = await session.execute(select(ChatContext).where(ChatContext.id == id))
|
||
|
|
data = result.scalar_one()
|
||
|
|
else:
|
||
|
|
result = await session.execute(select(ChatAnswer).where(ChatAnswer.id == id))
|
||
|
|
data = result.scalar_one()
|
||
|
|
await LearningChat.add_ban(data)
|
||
|
|
return {"status": 0, "msg": "禁用成功"}
|
||
|
|
except Exception as e:
|
||
|
|
return {"status": 500, "msg": f"禁用失败: {e}"}
|
||
|
|
|
||
|
|
@app.options("/learning_chat/api/ban_chat")
|
||
|
|
async def options_ban_chat():
|
||
|
|
return JSONResponse(content={}, status_code=200)
|
||
|
|
|
||
|
|
@app.put("/learning_chat/api/delete_all", response_class=JSONResponse, dependencies=[authentication()])
|
||
|
|
async def delete_all(type: str, id: Optional[int] = None):
|
||
|
|
try:
|
||
|
|
async with get_session() as session:
|
||
|
|
if type == "answer":
|
||
|
|
if id:
|
||
|
|
await session.execute(delete(ChatAnswer).where(ChatAnswer.context_id == id))
|
||
|
|
else:
|
||
|
|
await session.execute(delete(ChatAnswer))
|
||
|
|
elif type == "blacklist":
|
||
|
|
await session.execute(delete(ChatBlackList))
|
||
|
|
elif type == "context":
|
||
|
|
await session.execute(delete(ChatContext))
|
||
|
|
elif type == "message":
|
||
|
|
await session.execute(delete(ChatMessage))
|
||
|
|
await session.commit()
|
||
|
|
return {"status": 0, "msg": "操作成功"}
|
||
|
|
except Exception as e:
|
||
|
|
return {"status": 500, "msg": f"操作失败,{e}"}
|
||
|
|
|
||
|
|
@app.options("/learning_chat/api/delete_all")
|
||
|
|
async def options_delete_all():
|
||
|
|
return JSONResponse(content={}, status_code=200)
|
||
|
|
|
||
|
|
# 静态资源挂载放在所有 API 路由之后, 保证 /learning_chat/api/* 优先匹配
|
||
|
|
mount_frontend(app)
|