Add HeXi bot codebase: custom plugins, web frontends, tests
- hexi core: message handling, rate limiting, cooldown, plugin manager - Custom plugins: BF stats, daily check-in, quotes, persona cards, etc. - Community plugins vendored under hexi/plugins with local fixes - Web admin frontends (learning-chat, persona-admin), unified hexi/web - Tests for rate_limit/cooldown/memes/persona; poetry.lock Co-Authored-By: Claude <noreply@anthropic.com>
This commit is contained in:
@@ -0,0 +1,369 @@
|
||||
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)
|
||||
Reference in New Issue
Block a user