# -*- coding: utf-8 -*- """群聊学习 Web API 子应用(挂载到 /api/learning_chat)。 统一鉴权走 hexi.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_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