import datetime from typing import Optional, Union from fastapi import FastAPI from fastapi import Header, HTTPException, Depends from fastapi.responses import JSONResponse, HTMLResponse, RedirectResponse from jose import jwt from nonebot import get_app, get_adapter, logger from nonebot.adapters.onebot.v11 import Adapter from pydantic import BaseModel from sqlalchemy import select, delete, update, func try: import jieba_fast as jieba except ImportError: import jieba from nonebot_plugin_orm import get_session from .handler import LearningChat from .models import ChatMessage, ChatContext, ChatAnswer, ChatBlackList from .config import config_manager, driver from .web_frontend import ensure_frontend_ready, mount_frontend def authentication(): def inner(token: Optional[str] = Header(...)): try: payload = jwt.decode(token, config_manager.config.web_secret_key, algorithms="HS256") if ( not (username := payload.get("username")) or username != config_manager.config.web_username ): raise HTTPException(status_code=400, detail="登录验证失败或已失效,请重新登录") except (jwt.JWTError, jwt.ExpiredSignatureError, AttributeError): raise HTTPException(status_code=400, detail="登录验证失败或已失效,请重新登录") return Depends(inner) class UserModel(BaseModel): username: str password: str @driver.on_startup async def init_web(): if not config_manager.config.enable_web: return if not await ensure_frontend_ready(): logger.warning("群聊学习 | 前端构建失败, 管理页面将不可用") app: FastAPI = get_app() @app.post("/learning_chat/api/login", response_class=JSONResponse) async def login(user: UserModel): if ( user.username != config_manager.config.web_username or user.password != config_manager.config.web_password ): return {"status": -100, "msg": "登录失败,请确认用户ID和密码无误"} token = jwt.encode( { "username": user.username, "exp": datetime.datetime.now(datetime.timezone.utc) + datetime.timedelta(minutes=30), }, config_manager.config.web_secret_key, algorithm="HS256", ) return {"status": 0, "msg": "登录成功", "data": {"token": token}} @app.options("/learning_chat/api/login") async def options_login(): return JSONResponse(content={}, status_code=200) @app.get("/learning_chat/api/get_group_list", response_class=JSONResponse, dependencies=[authentication()]) async def get_group_list_api(): try: bots = get_adapter(Adapter).bots if len(bots) == 0: return {"status": -100, "msg": "获取群和好友列表失败,请确认已连接GOCQ"} bot = list(bots.values())[0] 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 {"status": 0, "msg": "ok", "data": {"group_list": group_list}} except ValueError: return {"status": -100, "msg": "获取群和好友列表失败,请确认已连接GOCQ"} @app.options("/learning_chat/api/get_group_list") async def options_get_group_list(): return JSONResponse(content={}, status_code=200) @app.get("/learning_chat/api/chat_global_config", response_class=JSONResponse, dependencies=[authentication()]) async def get_chat_global_config(): # 注意: 不再注入 member_list —— 前端全局配置页不使用(仅分群配置需要, # 由 chat_group_config 单独拉取单个群成员), 避免串行拉取全部群成员导致接口卡顿 return config_manager.config.dict(exclude={"group_config"}) @app.post("/learning_chat/api/chat_global_config", response_class=JSONResponse, dependencies=[authentication()]) 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.options("/learning_chat/api/chat_global_config") async def options_chat_global_config(): return JSONResponse(content={}, status_code=200) @app.get("/learning_chat/api/chat_group_config", response_class=JSONResponse, dependencies=[authentication()]) async def get_chat_group_config(group_id: int): try: bots = get_adapter(Adapter).bots if len(bots) == 0: return {"status": -100, "msg": "获取群和好友列表失败,请确认已连接GOCQ"} bot = list(bots.values())[0] 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 except ValueError: return {"status": -100, "msg": "获取群和好友列表失败,请确认已连接GOCQ"} @app.post("/learning_chat/api/chat_group_config", response_class=JSONResponse, dependencies=[authentication()]) async def post_chat_group_config(group_id: Union[int, str], data: dict): if not data.get("answer_threshold_weights"): return {"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 bots = get_adapter(Adapter).bots if len(bots) == 0: return {"status": -100, "msg": "获取群和好友列表失败,请确认已连接GOCQ"} bot = list(bots.values())[0] 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.options("/learning_chat/api/chat_group_config") async def options_chat_group_config(): return JSONResponse(content={}, status_code=200) @app.get("/learning_chat/api/get_chat_messages", response_class=JSONResponse, dependencies=[authentication()]) 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 {"status": 0, "msg": "ok", "data": {"items": items, "total": total}} @app.options("/learning_chat/api/get_chat_messages") async def options_get_chat_messages(): return JSONResponse(content={}, status_code=200) @app.get("/learning_chat/api/get_chat_contexts", response_class=JSONResponse, dependencies=[authentication()]) 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 {"status": 0, "msg": "ok", "data": {"items": items, "total": total}} @app.options("/learning_chat/api/get_chat_contexts") async def options_get_chat_contexts(): return JSONResponse(content={}, status_code=200) @app.get("/learning_chat/api/get_chat_answers", response_class=JSONResponse, dependencies=[authentication()]) 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"]] items.append(item) return {"status": 0, "msg": "ok", "data": {"items": items, "total": total}} @app.options("/learning_chat/api/get_chat_answers") async def options_get_chat_answers(): return JSONResponse(content={}, status_code=200) @app.get("/learning_chat/api/get_chat_blacklist", response_class=JSONResponse, dependencies=[authentication()]) 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() # bans 过滤需要基于全部数据(而非当前页), 故先全量取回再过滤、再分页 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 {"status": 0, "msg": "ok", "data": {"items": items, "total": total}} @app.options("/learning_chat/api/get_chat_blacklist") async def options_get_chat_blacklist(): return JSONResponse(content={}, status_code=200) @app.delete("/learning_chat/api/delete_chat", response_class=JSONResponse, dependencies=[authentication()]) 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 {"status": 500, "msg": f"删除失败,{e}"} @app.options("/learning_chat/api/delete_chat") async def options_delete_chat(): return JSONResponse(content={}, status_code=200) @app.put("/learning_chat/api/ban_chat", response_class=JSONResponse, dependencies=[authentication()]) 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 {"status": 500, "msg": f"禁用失败: {e}"} @app.options("/learning_chat/api/ban_chat") async def options_ban_chat(): return JSONResponse(content={}, status_code=200) @app.put("/learning_chat/api/delete_all", response_class=JSONResponse, dependencies=[authentication()]) 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 {"status": 500, "msg": f"操作失败,{e}"} @app.options("/learning_chat/api/delete_all") async def options_delete_all(): return JSONResponse(content={}, status_code=200) # 静态资源挂载放在所有 API 路由之后, 保证 /learning_chat/api/* 优先匹配 mount_frontend(app)