Files
HeXi/hexi/plugins/nonebot_plugin_learning_chat/web_hub.py
T
sansenhoshiandClaude b61d09f09f 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>
2026-09-01 13:13:40 +08:00

297 lines
13 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
# -*- coding: utf-8 -*-
"""群聊学习 Web API 子应用(挂载到 /api/learning_chat)。
统一鉴权走 hexi.web_auth(OAuth2 + SQLite),与统一管理台 /hub 共用登录态。
旧版独立后台 /learning_chat(JWT 自鉴权)保留可用,前端由 hub 渲染。
"""
from __future__ import annotations
from typing import Optional, Union
from fastapi import FastAPI
from fastapi.responses import JSONResponse
from sqlalchemy import delete, select, update, func
try:
import jieba_fast as jieba
except ImportError:
import jieba
from nonebot import get_adapter
from nonebot.adapters.onebot.v11 import Adapter
from nonebot_plugin_orm import get_session
from hexi.web_auth import require_admin
from .handler import LearningChat
from .models import ChatMessage, ChatContext, ChatAnswer, ChatBlackList
from .config import config_manager
API = require_admin
def _ok(data=None, msg: str = "ok") -> JSONResponse:
return JSONResponse({"status": 0, "msg": msg, "data": data})
async def _first_bot():
try:
bots = get_adapter(Adapter).bots
return next(iter(bots.values()), None)
except (ValueError, AttributeError):
return None
def build_admin_app() -> FastAPI | None:
"""构建群聊学习管理 API 子应用(挂载到 /api/learning_chat)。"""
if not config_manager.config.enable_web:
return None
app = FastAPI(title="Learning Chat API")
auth = API
@app.get("/get_group_list", response_class=JSONResponse, dependencies=[auth])
async def get_group_list_api():
bot = await _first_bot()
if bot is None:
return JSONResponse({"status": -100, "msg": "获取群和好友列表失败,请确认已连接GOCQ"})
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 _ok({"group_list": group_list})
@app.get("/chat_global_config", response_class=JSONResponse, dependencies=[auth])
async def get_chat_global_config():
return config_manager.config.dict(exclude={"group_config"})
@app.post("/chat_global_config", response_class=JSONResponse, dependencies=[auth])
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.get("/chat_group_config", response_class=JSONResponse, dependencies=[auth])
async def get_chat_group_config(group_id: int):
bot = await _first_bot()
if bot is None:
return JSONResponse({"status": -100, "msg": "获取群和好友列表失败,请确认已连接GOCQ"})
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
@app.post("/chat_group_config", response_class=JSONResponse, dependencies=[auth])
async def post_chat_group_config(group_id: Union[int, str], data: dict):
if not data.get("answer_threshold_weights"):
return JSONResponse({"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
bot = await _first_bot()
if bot is None:
return JSONResponse({"status": -100, "msg": "获取群和好友列表失败,请确认已连接GOCQ"})
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.get("/get_chat_messages", response_class=JSONResponse, dependencies=[auth])
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 _ok({"items": items, "total": total})
@app.get("/get_chat_contexts", response_class=JSONResponse, dependencies=[auth])
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 _ok({"items": items, "total": total})
@app.get("/get_chat_answers", response_class=JSONResponse, dependencies=[auth])
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"]] if item["messages"] else None
items.append(item)
return _ok({"items": items, "total": total})
@app.get("/get_chat_blacklist", response_class=JSONResponse, dependencies=[auth])
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()
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 _ok({"items": items, "total": total})
@app.delete("/delete_chat", response_class=JSONResponse, dependencies=[auth])
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 JSONResponse({"status": 500, "msg": f"删除失败,{e}"})
@app.put("/ban_chat", response_class=JSONResponse, dependencies=[auth])
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 JSONResponse({"status": 500, "msg": f"禁用失败: {e}"})
@app.put("/delete_all", response_class=JSONResponse, dependencies=[auth])
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 JSONResponse({"status": 500, "msg": f"操作失败,{e}"})
return app