2026-09-01 13:13:40 +08:00
|
|
|
|
# -*- 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
|
|
|
|
|
|
|
2026-09-03 00:44:38 +08:00
|
|
|
|
from .services.learn import LearningChat
|
2026-09-01 13:13:40 +08:00
|
|
|
|
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
|