294 lines
13 KiB
Python
294 lines
13 KiB
Python
# -*- coding: utf-8 -*-
|
||
"""群聊学习 Web API 子应用(挂载到 /api/learning_chat)。
|
||
|
||
统一鉴权走 hexi.web_hub.web_auth(OAuth2 + SQLite),与统一管理台 /hub 共用登录态。
|
||
前端由统一管理台 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
|
||
|
||
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
|