Files
HeXi/hexi/plugins/nonebot_plugin_learning_chat/web_hub.py
T

294 lines
13 KiB
Python
Raw Normal View History

# -*- coding: utf-8 -*-
"""群聊学习 Web API 子应用(挂载到 /api/learning_chat)。
2026-09-08 14:25:32 +08:00
统一鉴权走 hexi.web_hub.web_auth(OAuth2 + SQLite),与统一管理台 /hub 共用登录态。
2026-09-03 15:44:57 +08:00
前端由统一管理台 hexi/web 渲染。
"""
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
2026-09-08 14:25:32 +08:00
from hexi.web_hub.web_auth import require_admin
from .services.learn 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)。"""
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