Add HeXi bot codebase: custom plugins, web frontends, tests
- hexi core: message handling, rate limiting, cooldown, plugin manager - Custom plugins: BF stats, daily check-in, quotes, persona cards, etc. - Community plugins vendored under hexi/plugins with local fixes - Web admin frontends (learning-chat, persona-admin), unified hexi/web - Tests for rate_limit/cooldown/memes/persona; poetry.lock Co-Authored-By: Claude <noreply@anthropic.com>
@@ -0,0 +1,60 @@
|
||||
# nonebot_plugin_group_daily_analysis(OneBot V11 最小可用版)
|
||||
|
||||
移植自 astrbot_plugin_qq_group_daily_analysis(GitHub: SXP-Simon/astrbot_plugin_qq_group_daily_analysis),
|
||||
保留其领域分析 / 报告生成 / 持久化核心,把 AstrBot 全平台外壳重写为 NoneBot2 + OneBot V11。
|
||||
|
||||
## 群内指令(均需 @ 机器人,且群管理员 / 群主 / 超管)
|
||||
|
||||
- 群分析 [天数] :手动触发分析,可选最近 N 天(默认用配置里的时间窗口)
|
||||
- 查看模板 :列出可用报告模板及其编号,标出当前模板
|
||||
- 设置模板 [序号/名称] :切换报告模板,支持 1..9 编号或模板名
|
||||
- 分析设置 :查看当前分析参数(时间窗口 / 消息数 / 输出格式 / 模板 / 功能开关)
|
||||
- 设置分析 [参数] [值] :修改分析参数,见下表
|
||||
|
||||
设置分析 支持的参数:
|
||||
|
||||
天数/窗口 [N] 时间窗口,最近 N 天
|
||||
最大消息 [N] 单次最多拉取消息数
|
||||
最小消息 [N] 自动分析的最小消息阈值
|
||||
输出格式 [image/text]
|
||||
话题 / 称号 / 金句 / 聊天质量 [开/关]
|
||||
|
||||
## 已有能力
|
||||
|
||||
- 群分析:Onbot V11 get_group_msg_history 分页拉取近 N 天群聊
|
||||
- 话题分析 / 用户形象 / 群聊金句 / 聊天质量(OpenAI 兼容 LLM)
|
||||
- Playwright 渲染报告 HTML 为图片;报告内容以合并转发 (send_forward_msg) 发送,失败时回退直接发送
|
||||
- 历史摘要、本地 KV、断点检查点落盘到数据目录
|
||||
|
||||
## 尚未接入(后续阶段)
|
||||
|
||||
- 定时自动分析、增量分析、群漫画、多平台(Telegram/Discord/QQ 官方)
|
||||
- AstrBot 版 WebUI 看板
|
||||
- 群白名单管理指令
|
||||
|
||||
## 安装 / 启用
|
||||
|
||||
1. 插件目录已在 hexi/plugins/nonebot_plugin_group_daily_analysis,随 hexi/plugins 自动加载。
|
||||
2. 依赖:pip install ulid-py diskcache(已在 requirements.txt 追加)。
|
||||
3. 配置 LLM(环境变量,写入 NoneBot 的 .env):
|
||||
|
||||
HEXI_LLM_API_BASE=https://api.openai.com/v1
|
||||
HEXI_LLM_API_KEY=sk-xxx
|
||||
HEXI_LLM_MODEL=gpt-4o-mini
|
||||
|
||||
支持任何 OpenAI 兼容端点(One API / 中转 / 本地 vLLM 等)。
|
||||
|
||||
## 配置
|
||||
|
||||
配置文件:hexi/data/group_daily_analysis/config.json(首次运行由 _conf_schema.json 生成默认值)。
|
||||
命令里的 设置分析 / 设置模板 会直接写回该文件,改完即生效。
|
||||
|
||||
重点键:
|
||||
- basic.analysis_days:时间窗口(最近 N 天)
|
||||
- basic.max_messages:单次最大拉取消息数
|
||||
- basic.output_format:["image"] 或 ["text"]
|
||||
- basic.report_template:当前模板名(如 ATRI / scrapbook / BlueArchive)
|
||||
- analysis_features.topic_analysis_enabled / user_title_analysis_enabled / golden_quote_analysis_enabled / chat_quality_analysis_enabled
|
||||
- llm.llm_retries / llm.llm_backoff
|
||||
|
||||
注意:HEXI_LLM_API_BASE/KEY/MODEL 来自 .env 且优先于 config.json,不会被配置文件覆盖。
|
||||
@@ -0,0 +1,412 @@
|
||||
"""NoneBot 版「群日常分析」插件(OneBot V11 最小可用版)。
|
||||
|
||||
保留原 AstrBot 插件的领域分析/报告引擎,重写为 NoneBot 外壳:
|
||||
- on_command("群分析") 手动触发(需 @ 机器人)
|
||||
- OneBot V11 拉取群历史消息
|
||||
- OpenAI 兼容 LLM 通道
|
||||
- Playwright 渲染 HTML 报告为图片
|
||||
- 报告内容以合并转发 (send_forward_msg) 发送
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from nonebot import on_command, require
|
||||
from nonebot.adapters.onebot.v11 import Bot, GroupMessageEvent, MessageSegment
|
||||
from nonebot.log import logger
|
||||
from nonebot.params import CommandArg
|
||||
from nonebot.plugin import PluginMetadata
|
||||
from nonebot.rule import to_me
|
||||
|
||||
require("nonebot_plugin_alconna")
|
||||
from nonebot_plugin_alconna import UniMessage # noqa: E402
|
||||
|
||||
__plugin_meta__ = PluginMetadata(
|
||||
name="群分析插件",
|
||||
description="群日常分析总结插件(OneBot V11 最小可用版):生成群聊分析报告",
|
||||
usage="发送「群分析」或「群分析 [天数]」,需在群内 @ 机器人并由管理员触发",
|
||||
type="application",
|
||||
)
|
||||
|
||||
from .service import get_services, register_bot_adapter # noqa: E402
|
||||
from .renderer import html_render # noqa: E402
|
||||
from .templates import list_templates, template_exists # noqa: E402
|
||||
|
||||
|
||||
# ---------- 指令定义(全部要求 @ 机器人) ----------
|
||||
|
||||
analysis_cmd = on_command(
|
||||
"群分析",
|
||||
aliases={"group_analysis"},
|
||||
priority=5,
|
||||
block=True,
|
||||
rule=to_me(),
|
||||
)
|
||||
|
||||
view_templates_cmd = on_command(
|
||||
"查看模板",
|
||||
aliases={"模板列表", "view_templates"},
|
||||
priority=5,
|
||||
block=True,
|
||||
rule=to_me(),
|
||||
)
|
||||
|
||||
set_template_cmd = on_command(
|
||||
"设置模板",
|
||||
aliases={"set_template", "切换模板"},
|
||||
priority=5,
|
||||
block=True,
|
||||
rule=to_me(),
|
||||
)
|
||||
|
||||
analysis_settings_cmd = on_command(
|
||||
"分析设置",
|
||||
aliases={"查看参数", "analysis_settings"},
|
||||
priority=5,
|
||||
block=True,
|
||||
rule=to_me(),
|
||||
)
|
||||
|
||||
set_analysis_cmd = on_command(
|
||||
"设置分析",
|
||||
aliases={"set_analysis", "设置参数"},
|
||||
priority=5,
|
||||
block=True,
|
||||
rule=to_me(),
|
||||
)
|
||||
|
||||
|
||||
# ---------- 工具函数 ----------
|
||||
|
||||
def _is_privileged(event: GroupMessageEvent, bot: Bot) -> bool:
|
||||
from nonebot import get_driver
|
||||
try:
|
||||
superusers = getattr(get_driver().config, "superusers", [])
|
||||
if str(event.user_id) in [str(s) for s in superusers]:
|
||||
return True
|
||||
except Exception:
|
||||
pass
|
||||
role = getattr(event.sender, "role", "")
|
||||
return role in {"owner", "admin"}
|
||||
|
||||
|
||||
def _cmd_arg(a) -> str:
|
||||
"""把 CommandArg 里的参数转成纯字符串(可能是 MessageSegment)。"""
|
||||
return str(a).strip()
|
||||
|
||||
|
||||
def _cmd_tokens(args: tuple) -> list[str]:
|
||||
"""把 CommandArg 转成按空白切分的参数列表(兼容单段/多段两种返回)。"""
|
||||
tokens: list[str] = []
|
||||
for a in args:
|
||||
tokens.extend(_cmd_arg(a).split())
|
||||
return tokens
|
||||
|
||||
|
||||
def _parse_days(event: GroupMessageEvent) -> int | None:
|
||||
text = str(event.message)
|
||||
tokens = text.split()
|
||||
if len(tokens) >= 2:
|
||||
try:
|
||||
val = int(tokens[1])
|
||||
return max(1, min(val, 30)) if val > 0 else None
|
||||
except ValueError:
|
||||
return None
|
||||
return None
|
||||
|
||||
|
||||
def _build_text_report(result: dict) -> str:
|
||||
ar = result.get("analysis_result", {})
|
||||
stats = ar.get("statistics")
|
||||
lines = ["[群聊分析报告(文本版)]"]
|
||||
if stats:
|
||||
mc = getattr(stats, "message_count", "?")
|
||||
pc = getattr(stats, "participant_count", "?")
|
||||
lines.append(f"发言 {mc} 条 / 参与 {pc} 人")
|
||||
topics = ar.get("topics", [])
|
||||
if topics:
|
||||
lines.append("")
|
||||
lines.append("[今日话题]")
|
||||
for i, t in enumerate(topics[:6], 1):
|
||||
title = getattr(t, "topic", "") or (t.get("topic") if isinstance(t, dict) else "")
|
||||
detail = getattr(t, "detail", "") or (t.get("detail") if isinstance(t, dict) else "")
|
||||
lines.append(f"{i}. {title}: {detail}")
|
||||
quotes = ar.get("golden_quotes", [])
|
||||
if quotes:
|
||||
lines.append("")
|
||||
lines.append("[群聊圣经]")
|
||||
for q in quotes[:5]:
|
||||
text = getattr(q, "quote", "") or (q.get("quote") if isinstance(q, dict) else "")
|
||||
user = getattr(q, "user_name", "") or (q.get("user_name") if isinstance(q, dict) else "")
|
||||
lines.append(f"[{text}] —— {user}")
|
||||
lines.append("")
|
||||
lines.append("(文本版,图片报告请将输出格式设为 image)")
|
||||
return chr(10).join(lines)
|
||||
|
||||
|
||||
def _build_forward_nodes(adapter, group_id: str, text_report: str, image_url: str | None) -> list[dict]:
|
||||
"""构建 OneBot v11 合并转发 (send_forward_msg) 节点列表。"""
|
||||
self_uin = (adapter.bot_self_ids[0] if adapter and adapter.bot_self_ids else "10000")
|
||||
nodes = []
|
||||
nodes.append(
|
||||
{
|
||||
"type": "node",
|
||||
"data": {
|
||||
"name": "群聊分析报告",
|
||||
"uin": self_uin,
|
||||
"content": [MessageSegment.text(f"群聊分析报告(群 {group_id})")],
|
||||
},
|
||||
}
|
||||
)
|
||||
if text_report:
|
||||
nodes.append(
|
||||
{
|
||||
"type": "node",
|
||||
"data": {
|
||||
"name": "报告内容",
|
||||
"uin": self_uin,
|
||||
"content": [MessageSegment.text(text_report)],
|
||||
},
|
||||
}
|
||||
)
|
||||
if image_url:
|
||||
nodes.append(
|
||||
{
|
||||
"type": "node",
|
||||
"data": {
|
||||
"name": "报告图片",
|
||||
"uin": self_uin,
|
||||
"content": [MessageSegment.image(file=image_url)],
|
||||
},
|
||||
}
|
||||
)
|
||||
return nodes
|
||||
|
||||
|
||||
# ---------- 手动分析 ----------
|
||||
|
||||
@analysis_cmd.handle()
|
||||
async def _(bot: Bot, event: GroupMessageEvent):
|
||||
if not _is_privileged(event, bot):
|
||||
await UniMessage.text("仅群管理员或超级管理员可使用该指令").send()
|
||||
return
|
||||
|
||||
group_id = str(event.group_id)
|
||||
days = _parse_days(event)
|
||||
await UniMessage.text("正在启动分析引擎,正在拉取最近消息...").send()
|
||||
|
||||
try:
|
||||
svc = get_services()
|
||||
register_bot_adapter(bot)
|
||||
analysis_service = svc["analysis_service"]
|
||||
report_generator = svc["report_generator"]
|
||||
config_manager = svc["config_manager"]
|
||||
|
||||
result = await analysis_service.execute_daily_analysis(
|
||||
group_id=group_id,
|
||||
platform_id="onebot",
|
||||
manual=True,
|
||||
days=days,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f"群分析异常: {e}", exc_info=True)
|
||||
await UniMessage.text(f"分析失败: {e}").send()
|
||||
return
|
||||
|
||||
if not result.get("success"):
|
||||
reason = result.get("reason")
|
||||
if reason == "no_messages":
|
||||
await UniMessage.text("未找到足够的群聊记录").send()
|
||||
elif reason == "llm_analysis_failed":
|
||||
await UniMessage.text("大模型文本分析失败(请检查 LLM API Key 及服务商连通性)").send()
|
||||
else:
|
||||
await UniMessage.text(f"分析失败: {result.get('error', '原因未知')}").send()
|
||||
return
|
||||
|
||||
analysis_result = result["analysis_result"]
|
||||
adapter = result.get("adapter")
|
||||
platform_id = result.get("platform_id", "onebot")
|
||||
|
||||
async def avatar_url_getter(user_id: str) -> str | None:
|
||||
if adapter:
|
||||
return await adapter.get_user_avatar_url(user_id)
|
||||
return None
|
||||
|
||||
async def nickname_getter(user_id: str) -> str | None:
|
||||
if adapter:
|
||||
try:
|
||||
member = await adapter.get_member_info(group_id, user_id)
|
||||
if member:
|
||||
return member.card or member.nickname
|
||||
except Exception:
|
||||
pass
|
||||
return None
|
||||
|
||||
output_format = config_manager.get_output_format()
|
||||
fmt = str(output_format[0]).lower() if output_format else "image"
|
||||
text_report = _build_text_report(result) or "(无文本报告)"
|
||||
image_url = None
|
||||
|
||||
if fmt != "text":
|
||||
await UniMessage.text("正在生成报告图片...").send()
|
||||
try:
|
||||
image_url, html_content = await report_generator.generate_image_report(
|
||||
analysis_result,
|
||||
group_id,
|
||||
html_render,
|
||||
avatar_url_getter=avatar_url_getter,
|
||||
nickname_getter=nickname_getter,
|
||||
avatar_cache_namespace=platform_id,
|
||||
allow_alphanumeric_user_ids=False,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f"报告生成异常: {e}", exc_info=True)
|
||||
image_url = None
|
||||
|
||||
# 用合并转发发送报告内容
|
||||
nodes = _build_forward_nodes(adapter, group_id, text_report, image_url)
|
||||
sent = False
|
||||
if adapter:
|
||||
try:
|
||||
sent = await adapter.send_forward_msg(group_id, nodes)
|
||||
except Exception as e:
|
||||
logger.warning(f"发送合并转发失败: {e}")
|
||||
sent = False
|
||||
|
||||
# 合并转发失败 -> 回退直接发送
|
||||
if not sent:
|
||||
try:
|
||||
if image_url and adapter:
|
||||
sent = await adapter.send_image(group_id, image_url)
|
||||
if not sent:
|
||||
await UniMessage.text(text_report).send()
|
||||
except Exception as e:
|
||||
logger.warning(f"回退发送失败: {e}")
|
||||
await UniMessage.text(text_report).send()
|
||||
|
||||
|
||||
# ---------- 模板 / 参数配置指令 ----------
|
||||
|
||||
@view_templates_cmd.handle()
|
||||
async def _(bot: Bot, event: GroupMessageEvent):
|
||||
if not _is_privileged(event, bot):
|
||||
await UniMessage.text("仅群管理员或超级管理员可使用该指令").send()
|
||||
return
|
||||
svc = get_services()
|
||||
cm = svc["config_manager"]
|
||||
current = cm.get_report_template()
|
||||
names = list_templates()
|
||||
lines = ["可用报告模板列表", f"当前使用: {current}", "使用「设置模板 序号」切换"]
|
||||
for i, name in enumerate(names, 1):
|
||||
mark = " [当前]" if name == current else ""
|
||||
lines.append(f"{i}. {name}{mark}")
|
||||
await UniMessage.text(chr(10).join(lines)).send()
|
||||
|
||||
|
||||
@set_template_cmd.handle()
|
||||
async def _(bot: Bot, event: GroupMessageEvent, args: tuple = CommandArg()):
|
||||
if not _is_privileged(event, bot):
|
||||
await UniMessage.text("仅群管理员或超级管理员可使用该指令").send()
|
||||
return
|
||||
tokens = _cmd_tokens(args)
|
||||
arg = tokens[0] if tokens else ""
|
||||
names = list_templates()
|
||||
target = None
|
||||
if arg.isdigit():
|
||||
idx = int(arg)
|
||||
if 1 <= idx <= len(names):
|
||||
target = names[idx - 1]
|
||||
else:
|
||||
await UniMessage.text(f"序号无效,有效范围 1-{len(names)}").send()
|
||||
return
|
||||
elif template_exists(arg):
|
||||
target = arg
|
||||
else:
|
||||
await UniMessage.text(f"模板不存在,可用:{'、'.join(names)}").send()
|
||||
return
|
||||
svc = get_services()
|
||||
cm = svc["config_manager"]
|
||||
cm.set_report_template(target)
|
||||
await UniMessage.text(f"已设置模板:{target}(编号 {names.index(target) + 1})").send()
|
||||
|
||||
|
||||
@analysis_settings_cmd.handle()
|
||||
async def _(bot: Bot, event: GroupMessageEvent):
|
||||
if not _is_privileged(event, bot):
|
||||
await UniMessage.text("仅群管理员或超级管理员可使用该指令").send()
|
||||
return
|
||||
cm = get_services()["config_manager"]
|
||||
fmt = "、".join(str(x) for x in cm.get_output_format())
|
||||
lines = [
|
||||
"当前分析参数",
|
||||
f"时间窗口: {cm.get_analysis_days()} 天",
|
||||
f"最大消息数: {cm.get_max_messages()}",
|
||||
f"最小消息数: {cm.get_min_messages_threshold()}",
|
||||
f"输出格式: {fmt}",
|
||||
f"报告模板: {cm.get_report_template()}",
|
||||
f"话题分析: {'开' if cm.get_topic_analysis_enabled() else '关'}",
|
||||
f"用户称号: {'开' if cm.get_user_title_analysis_enabled() else '关'}",
|
||||
f"金句分析: {'开' if cm.get_golden_quote_analysis_enabled() else '关'}",
|
||||
f"聊天质量: {'开' if cm.get_chat_quality_analysis_enabled() else '关'}",
|
||||
"用法:设置分析 [参数] [值],如:设置分析 天数 3",
|
||||
]
|
||||
await UniMessage.text(chr(10).join(lines)).send()
|
||||
|
||||
|
||||
_SETTING_MAP = {
|
||||
"天数": ("set_analysis_days", int, "分析天数已设为 {} 天"),
|
||||
"窗口": ("set_analysis_days", int, "分析窗口已设为 {} 天"),
|
||||
"最大消息": ("set_max_messages", int, "最大消息数已设为 {}"),
|
||||
"最小消息": ("set_min_messages_threshold", int, "最小消息数已设为 {}"),
|
||||
"输出格式": ("set_output_format", str, "输出格式已设为 {}"),
|
||||
"话题": ("set_topic_analysis_enabled", str, "话题分析已{}"),
|
||||
"称号": ("set_user_title_analysis_enabled", str, "用户称号已{}"),
|
||||
"金句": ("set_golden_quote_analysis_enabled", str, "金句分析已{}"),
|
||||
"聊天质量": ("set_chat_quality_analysis_enabled", str, "聊天质量已{}"),
|
||||
}
|
||||
|
||||
_BOOL_WORDS = {"开": True, "on": True, "true": True, "关": False, "off": False, "false": False}
|
||||
|
||||
|
||||
@set_analysis_cmd.handle()
|
||||
async def _(bot: Bot, event: GroupMessageEvent, args: tuple = CommandArg()):
|
||||
if not _is_privileged(event, bot):
|
||||
await UniMessage.text("仅群管理员或超级管理员可使用该指令").send()
|
||||
return
|
||||
tokens = _cmd_tokens(args)
|
||||
if len(tokens) < 2:
|
||||
await UniMessage.text(
|
||||
"用法:设置分析 [参数] [值]" + chr(10) + "参数:天数/窗口/最大消息/最小消息/输出格式/话题/称号/金句/聊天质量"
|
||||
).send()
|
||||
return
|
||||
key = tokens[0]
|
||||
val = tokens[1]
|
||||
mapping = _SETTING_MAP.get(key)
|
||||
if not mapping:
|
||||
await UniMessage.text(f"未知参数:{key};可用:{'、'.join(_SETTING_MAP)}").send()
|
||||
return
|
||||
setter_name, conv, _ = mapping
|
||||
cm = get_services()["config_manager"]
|
||||
setter = getattr(cm, setter_name)
|
||||
try:
|
||||
if conv is int:
|
||||
v = int(val)
|
||||
if v < 0:
|
||||
raise ValueError("不能为负数")
|
||||
setter(v)
|
||||
await UniMessage.text(f"{key} 已设为 {v}").send()
|
||||
elif conv is str:
|
||||
v = val.lower()
|
||||
if key in ("话题", "称号", "金句", "聊天质量"):
|
||||
if v not in _BOOL_WORDS:
|
||||
await UniMessage.text("布尔值请填:开/关 或 on/off").send()
|
||||
return
|
||||
setter(_BOOL_WORDS[v])
|
||||
await UniMessage.text(f"{key} 已{'开启' if _BOOL_WORDS[v] else '关闭'}").send()
|
||||
else:
|
||||
if v not in ("image", "text"):
|
||||
await UniMessage.text("输出格式仅支持 image/text").send()
|
||||
return
|
||||
setter(["text"] if v == "text" else ["image"])
|
||||
await UniMessage.text(f"输出格式已设为 {v}").send()
|
||||
except (TypeError, ValueError) as e:
|
||||
await UniMessage.text(f"参数值无效: {e}").send()
|
||||
@@ -0,0 +1,307 @@
|
||||
"""NoneBot OneBot V11 适配器:负责拉取群历史消息、群组/成员信息与发送。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import base64
|
||||
import time
|
||||
from datetime import datetime, timedelta
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from .core.domain.value_objects.unified_group import UnifiedGroup, UnifiedMember
|
||||
from .core.domain.value_objects.unified_message import (
|
||||
MessageContent,
|
||||
MessageContentType,
|
||||
UnifiedMessage,
|
||||
)
|
||||
from .core.utils.logger import logger
|
||||
|
||||
|
||||
class OneBotAdapter:
|
||||
"""面向 NoneBot OneBot V11 的最小适配器。"""
|
||||
|
||||
platform_name = "onebot"
|
||||
|
||||
USER_AVATAR_TEMPLATE = "https://q1.qlogo.cn/g?b=qq&nk={user_id}&s=640"
|
||||
|
||||
def __init__(self, bot: Any, config: dict | None = None):
|
||||
self.bot = bot
|
||||
self.config = config or {}
|
||||
self.platform_id = str(self.config.get("platform_id") or "onebot")
|
||||
self.bot_self_ids = [str(x) for x in self.config.get("bot_self_ids", [])]
|
||||
self.filter_bot_messages = bool(self.config.get("filter_bot_messages", True))
|
||||
|
||||
# —— 消息拉取 ——
|
||||
async def fetch_messages(
|
||||
self,
|
||||
group_id: str,
|
||||
days: int = 1,
|
||||
max_count: int = 1000,
|
||||
before_id: str | None = None,
|
||||
since_ts: int | None = None,
|
||||
) -> list[UnifiedMessage]:
|
||||
try:
|
||||
chunk_size = 100
|
||||
all_raw = []
|
||||
seen_raw_ids: set[str] = set()
|
||||
|
||||
if since_ts and since_ts > 0:
|
||||
start_ts = since_ts
|
||||
else:
|
||||
start_ts = int(
|
||||
(datetime.now() - timedelta(days=days)).timestamp()
|
||||
)
|
||||
|
||||
current_anchor = before_id
|
||||
while len(all_raw) < max_count:
|
||||
fetch_count = min(chunk_size, max_count - len(all_raw))
|
||||
params: dict[str, Any] = {
|
||||
"group_id": int(group_id),
|
||||
"count": fetch_count,
|
||||
}
|
||||
if current_anchor:
|
||||
params["message_seq"] = current_anchor
|
||||
params["reverseOrder"] = True
|
||||
|
||||
result = None
|
||||
for attempt in range(1, 4):
|
||||
try:
|
||||
result = await self.bot.call_api(
|
||||
"get_group_msg_history", **params
|
||||
)
|
||||
break
|
||||
except Exception as exc:
|
||||
if attempt < 3:
|
||||
logger.warning(
|
||||
f"OneBot 分页拉取失败(第{attempt}次): {exc}"
|
||||
)
|
||||
await asyncio.sleep(attempt)
|
||||
else:
|
||||
logger.warning(
|
||||
f"OneBot 分页拉取重试耗尽: {exc}"
|
||||
)
|
||||
if not result or "messages" not in result:
|
||||
break
|
||||
messages = result.get("messages", [])
|
||||
if not messages:
|
||||
break
|
||||
|
||||
first = messages[0]
|
||||
last = messages[-1]
|
||||
earliest = first if first.get("time", 0) <= last.get("time", 0) else last
|
||||
chunk_earliest_ts = earliest.get("time", 0)
|
||||
|
||||
for raw in messages:
|
||||
msg_time = raw.get("time", 0)
|
||||
msg_id = str(raw.get("message_id", ""))
|
||||
if not msg_id or msg_id in seen_raw_ids:
|
||||
continue
|
||||
if start_ts <= msg_time <= int(datetime.now().timestamp()):
|
||||
all_raw.append(raw)
|
||||
seen_raw_ids.add(msg_id)
|
||||
|
||||
seq_val = (
|
||||
earliest.get("message_seq")
|
||||
or earliest.get("real_id")
|
||||
or earliest.get("seq")
|
||||
)
|
||||
mid_val = earliest.get("message_id")
|
||||
new_anchor = seq_val if seq_val is not None else mid_val
|
||||
if chunk_earliest_ts <= start_ts:
|
||||
break
|
||||
if current_anchor and str(new_anchor) == str(current_anchor):
|
||||
break
|
||||
current_anchor = new_anchor
|
||||
await asyncio.sleep(0.05)
|
||||
|
||||
unified: list[UnifiedMessage] = []
|
||||
seen: set[str] = set()
|
||||
for raw in all_raw:
|
||||
mid = str(raw.get("message_id", ""))
|
||||
if not mid or mid in seen:
|
||||
continue
|
||||
u = self._convert_message(raw, group_id)
|
||||
if u:
|
||||
unified.append(u)
|
||||
seen.add(mid)
|
||||
unified.sort(key=lambda m: m.timestamp)
|
||||
return unified
|
||||
except Exception as e:
|
||||
logger.warning(f"OneBot 分页获取消息失败: {e}")
|
||||
return []
|
||||
|
||||
def _convert_message(self, raw: dict, group_id: str) -> UnifiedMessage | None:
|
||||
try:
|
||||
sender = raw.get("sender", {})
|
||||
chain = raw.get("message", [])
|
||||
if isinstance(chain, str):
|
||||
chain = [{"type": "text", "data": {"text": chain}}]
|
||||
contents: list[MessageContent] = []
|
||||
text_parts: list[str] = []
|
||||
for seg in chain:
|
||||
seg_t = seg.get("type", "")
|
||||
seg_d = seg.get("data", {})
|
||||
if seg_t == "text":
|
||||
text = seg_d.get("text", "")
|
||||
text_parts.append(text)
|
||||
contents.append(MessageContent(type=MessageContentType.TEXT, text=text))
|
||||
elif seg_t == "image":
|
||||
sub_type = seg_d.get("subType", seg_d.get("sub_type"))
|
||||
try:
|
||||
is_sticker = int(sub_type) == 1
|
||||
except (TypeError, ValueError):
|
||||
is_sticker = False
|
||||
raw_data: dict[str, Any] = {"summary": seg_d.get("summary", "")}
|
||||
if sub_type is not None:
|
||||
try:
|
||||
raw_data["sub_type"] = int(sub_type)
|
||||
except (TypeError, ValueError):
|
||||
pass
|
||||
contents.append(
|
||||
MessageContent(
|
||||
type=MessageContentType.EMOJI if is_sticker else MessageContentType.IMAGE,
|
||||
url=seg_d.get("url", seg_d.get("file", "")),
|
||||
raw_data=raw_data,
|
||||
)
|
||||
)
|
||||
elif seg_t == "at":
|
||||
contents.append(
|
||||
MessageContent(type=MessageContentType.AT, at_user_id=str(seg_d.get("qq", "")))
|
||||
)
|
||||
elif seg_t in ("face", "mface", "bface", "sface"):
|
||||
contents.append(
|
||||
MessageContent(
|
||||
type=MessageContentType.EMOJI,
|
||||
emoji_id=str(seg_d.get("id", "")),
|
||||
raw_data={"face_type": seg_t},
|
||||
)
|
||||
)
|
||||
elif seg_t == "reply":
|
||||
contents.append(
|
||||
MessageContent(type=MessageContentType.REPLY, raw_data={"reply_id": seg_d.get("id", "")})
|
||||
)
|
||||
else:
|
||||
contents.append(MessageContent(type=MessageContentType.UNKNOWN, raw_data=seg))
|
||||
|
||||
reply_to = None
|
||||
for c in contents:
|
||||
if c.type == MessageContentType.REPLY and c.raw_data:
|
||||
reply_to = str(c.raw_data.get("reply_id", ""))
|
||||
break
|
||||
|
||||
return UnifiedMessage(
|
||||
message_id=str(raw.get("message_id", "")),
|
||||
sender_id=str(sender.get("user_id", "")),
|
||||
sender_name=sender.get("nickname", ""),
|
||||
sender_card=sender.get("card", "") or None,
|
||||
group_id=group_id,
|
||||
text_content="".join(text_parts),
|
||||
contents=tuple(contents),
|
||||
timestamp=raw.get("time", 0),
|
||||
platform="onebot",
|
||||
reply_to_id=reply_to,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.debug(f"OneBot _convert_message error: {e}")
|
||||
return None
|
||||
|
||||
# —— 群/成员信息 ——
|
||||
async def get_group_info(self, group_id: str) -> UnifiedGroup | None:
|
||||
try:
|
||||
gid = int(group_id)
|
||||
data = await self.bot.call_api("get_group_info", group_id=gid)
|
||||
if isinstance(data, list):
|
||||
data = next((d for d in data if str(d.get("group_id")) == str(gid)), data[0] if data else {})
|
||||
return UnifiedGroup(
|
||||
group_id=str(gid),
|
||||
group_name=data.get("group_name", ""),
|
||||
member_count=int(data.get("member_count", 0) or 0),
|
||||
platform="onebot",
|
||||
)
|
||||
except Exception as e:
|
||||
logger.debug(f"get_group_info({group_id}) 失败: {e}")
|
||||
return None
|
||||
|
||||
async def get_user_avatar_url(self, user_id: str) -> str | None:
|
||||
return self.USER_AVATAR_TEMPLATE.format(user_id=user_id)
|
||||
|
||||
async def get_member_info(self, group_id: str, user_id: str) -> UnifiedMember | None:
|
||||
try:
|
||||
data = await self.bot.call_api(
|
||||
"get_group_member_info",
|
||||
group_id=int(group_id),
|
||||
user_id=int(user_id),
|
||||
)
|
||||
return UnifiedMember(
|
||||
user_id=str(user_id),
|
||||
nickname=data.get("nickname", ""),
|
||||
card=data.get("card", "") or None,
|
||||
role=str(data.get("role", "member")),
|
||||
)
|
||||
except Exception as e:
|
||||
logger.debug(f"get_member_info({group_id},{user_id}) 失败: {e}")
|
||||
return None
|
||||
|
||||
async def is_group_muted(self, group_id: str) -> bool:
|
||||
return False
|
||||
|
||||
def get_platform_name(self) -> str:
|
||||
return self.platform_name
|
||||
|
||||
# —— 发送(在手动/自动报告中可用;手动指令会走 SAA 发送) ——
|
||||
async def send_text(self, group_id: str, text: str) -> bool:
|
||||
try:
|
||||
await self.bot.call_api("send_group_msg", group_id=int(group_id), message=str(text))
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.warning(f"send_text 失败: {e}")
|
||||
return False
|
||||
|
||||
async def send_image(self, group_id: str, image_url: str, caption: str = "") -> bool:
|
||||
try:
|
||||
if image_url.startswith("base64://"):
|
||||
b64 = image_url.split("base64://", 1)[1]
|
||||
from nonebot.adapters.onebot.v11 import MessageSegment
|
||||
msg = MessageSegment.image(file=f"base64://{b64}")
|
||||
else:
|
||||
from nonebot.adapters.onebot.v11 import MessageSegment
|
||||
msg = MessageSegment.image(file=image_url)
|
||||
await self.bot.call_api("send_group_msg", group_id=int(group_id), message=msg)
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.warning(f"send_image 失败: {e}")
|
||||
return False
|
||||
|
||||
async def send_file(self, group_id: str, file_path: str, caption: str = "") -> bool:
|
||||
return False
|
||||
|
||||
async def send_forward_msg(self, group_id: str, nodes: list[dict]) -> bool:
|
||||
"""发送 OneBot v11 合并转发消息 (send_forward_msg)。"""
|
||||
try:
|
||||
from nonebot.adapters.onebot.v11 import MessageSegment
|
||||
|
||||
normalized = []
|
||||
for n in nodes:
|
||||
data = dict(n.get("data", {}))
|
||||
content = data.get("content")
|
||||
if isinstance(content, str):
|
||||
content = [MessageSegment.text(content)]
|
||||
elif content is None:
|
||||
content = []
|
||||
data["content"] = content
|
||||
normalized.append({"type": "node", "data": data})
|
||||
|
||||
await self.bot.call_api(
|
||||
"send_forward_msg", group_id=int(group_id), messages=normalized
|
||||
)
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.warning(f"send_forward_msg 失败: {e}")
|
||||
return False
|
||||
|
||||
async def set_reaction(self, group_id: str, message_id: str, emoji: str | int, is_add: bool = True) -> bool:
|
||||
return False
|
||||
|
||||
async def prepare_group_member_cache(self, group_id: str):
|
||||
return True, None
|
||||
|
After Width: | Height: | Size: 540 KiB |
|
After Width: | Height: | Size: 726 KiB |
|
After Width: | Height: | Size: 690 KiB |
|
After Width: | Height: | Size: 1.7 MiB |
|
After Width: | Height: | Size: 913 KiB |
|
After Width: | Height: | Size: 793 KiB |
|
After Width: | Height: | Size: 951 KiB |
|
After Width: | Height: | Size: 1.1 MiB |
|
After Width: | Height: | Size: 630 KiB |
|
After Width: | Height: | Size: 1.2 MiB |
|
After Width: | Height: | Size: 1.3 MiB |
|
After Width: | Height: | Size: 271 KiB |
|
After Width: | Height: | Size: 677 KiB |
|
After Width: | Height: | Size: 652 KiB |
@@ -0,0 +1,246 @@
|
||||
{
|
||||
"sbti": [
|
||||
{
|
||||
"code": "ATM-er",
|
||||
"file": "https://fastly.jsdelivr.net/gh/SXP-Simon/profile_assets@main/sbti/ATM-er.png"
|
||||
},
|
||||
{
|
||||
"code": "BOSS",
|
||||
"file": "https://fastly.jsdelivr.net/gh/SXP-Simon/profile_assets@main/sbti/BOSS.png"
|
||||
},
|
||||
{
|
||||
"code": "CTRL",
|
||||
"file": "https://fastly.jsdelivr.net/gh/SXP-Simon/profile_assets@main/sbti/CTRL.png"
|
||||
},
|
||||
{
|
||||
"code": "DEAD",
|
||||
"file": "https://fastly.jsdelivr.net/gh/SXP-Simon/profile_assets@main/sbti/DEAD.png"
|
||||
},
|
||||
{
|
||||
"code": "Dior-s",
|
||||
"file": "https://fastly.jsdelivr.net/gh/SXP-Simon/profile_assets@main/sbti/Dior-s.jpg"
|
||||
},
|
||||
{
|
||||
"code": "DRUNK",
|
||||
"file": "https://fastly.jsdelivr.net/gh/SXP-Simon/profile_assets@main/sbti/DRUNK.png"
|
||||
},
|
||||
{
|
||||
"code": "FAKE",
|
||||
"file": "https://fastly.jsdelivr.net/gh/SXP-Simon/profile_assets@main/sbti/FAKE.png"
|
||||
},
|
||||
{
|
||||
"code": "FUCK",
|
||||
"file": "https://fastly.jsdelivr.net/gh/SXP-Simon/profile_assets@main/sbti/FUCK.png"
|
||||
},
|
||||
{
|
||||
"code": "GOGO",
|
||||
"file": "https://fastly.jsdelivr.net/gh/SXP-Simon/profile_assets@main/sbti/GOGO.png"
|
||||
},
|
||||
{
|
||||
"code": "HHHH",
|
||||
"file": "https://fastly.jsdelivr.net/gh/SXP-Simon/profile_assets@main/sbti/HHHH.png"
|
||||
},
|
||||
{
|
||||
"code": "IMFW",
|
||||
"file": "https://fastly.jsdelivr.net/gh/SXP-Simon/profile_assets@main/sbti/IMFW.png"
|
||||
},
|
||||
{
|
||||
"code": "IMSB",
|
||||
"file": "https://fastly.jsdelivr.net/gh/SXP-Simon/profile_assets@main/sbti/IMSB.png"
|
||||
},
|
||||
{
|
||||
"code": "JOKE-R",
|
||||
"file": "https://fastly.jsdelivr.net/gh/SXP-Simon/profile_assets@main/sbti/JOKE-R.jpg"
|
||||
},
|
||||
{
|
||||
"code": "LOVE-R",
|
||||
"file": "https://fastly.jsdelivr.net/gh/SXP-Simon/profile_assets@main/sbti/LOVE-R.png"
|
||||
},
|
||||
{
|
||||
"code": "MALO",
|
||||
"file": "https://fastly.jsdelivr.net/gh/SXP-Simon/profile_assets@main/sbti/MALO.png"
|
||||
},
|
||||
{
|
||||
"code": "MONK",
|
||||
"file": "https://fastly.jsdelivr.net/gh/SXP-Simon/profile_assets@main/sbti/MONK.png"
|
||||
},
|
||||
{
|
||||
"code": "MUM",
|
||||
"file": "https://fastly.jsdelivr.net/gh/SXP-Simon/profile_assets@main/sbti/MUM.png"
|
||||
},
|
||||
{
|
||||
"code": "OH-NO",
|
||||
"file": "https://fastly.jsdelivr.net/gh/SXP-Simon/profile_assets@main/sbti/OH-NO.png"
|
||||
},
|
||||
{
|
||||
"code": "OJBK",
|
||||
"file": "https://fastly.jsdelivr.net/gh/SXP-Simon/profile_assets@main/sbti/OJBK.png"
|
||||
},
|
||||
{
|
||||
"code": "POOR",
|
||||
"file": "https://fastly.jsdelivr.net/gh/SXP-Simon/profile_assets@main/sbti/POOR.png"
|
||||
},
|
||||
{
|
||||
"code": "SEXY",
|
||||
"file": "https://fastly.jsdelivr.net/gh/SXP-Simon/profile_assets@main/sbti/SEXY.png"
|
||||
},
|
||||
{
|
||||
"code": "SHIT",
|
||||
"file": "https://fastly.jsdelivr.net/gh/SXP-Simon/profile_assets@main/sbti/SHIT.png"
|
||||
},
|
||||
{
|
||||
"code": "SOLO",
|
||||
"file": "https://fastly.jsdelivr.net/gh/SXP-Simon/profile_assets@main/sbti/SOLO.png"
|
||||
},
|
||||
{
|
||||
"code": "THAN-K",
|
||||
"file": "https://fastly.jsdelivr.net/gh/SXP-Simon/profile_assets@main/sbti/THAN-K.png"
|
||||
},
|
||||
{
|
||||
"code": "THIN-K",
|
||||
"file": "https://fastly.jsdelivr.net/gh/SXP-Simon/profile_assets@main/sbti/THIN-K.png"
|
||||
},
|
||||
{
|
||||
"code": "WOC",
|
||||
"file": "https://fastly.jsdelivr.net/gh/SXP-Simon/profile_assets@main/sbti/WOC.png"
|
||||
},
|
||||
{
|
||||
"code": "ZZZZ",
|
||||
"file": "https://fastly.jsdelivr.net/gh/SXP-Simon/profile_assets@main/sbti/ZZZZ.png"
|
||||
}
|
||||
],
|
||||
"acgti": [
|
||||
{
|
||||
"code": "CIRN",
|
||||
"name": "琪露诺",
|
||||
"mbti": "ESTP",
|
||||
"file": "https://fastly.jsdelivr.net/gh/SXP-Simon/profile_assets@main/acgti/characters/CIRN.png"
|
||||
},
|
||||
{
|
||||
"code": "MIKT",
|
||||
"name": "御坂美琴",
|
||||
"mbti": "ESTJ",
|
||||
"file": "https://fastly.jsdelivr.net/gh/SXP-Simon/profile_assets@main/acgti/characters/MIKT.png"
|
||||
},
|
||||
{
|
||||
"code": "MADK",
|
||||
"name": "鹿目圆",
|
||||
"mbti": "INFP",
|
||||
"file": "https://fastly.jsdelivr.net/gh/SXP-Simon/profile_assets@main/acgti/characters/MADK.png"
|
||||
},
|
||||
{
|
||||
"code": "FRNA",
|
||||
"name": "芙宁娜",
|
||||
"mbti": "ESFP",
|
||||
"file": "https://fastly.jsdelivr.net/gh/SXP-Simon/profile_assets@main/acgti/characters/FRNA.png"
|
||||
},
|
||||
{
|
||||
"code": "GGGG-A",
|
||||
"name": "高松灯",
|
||||
"mbti": "INFP",
|
||||
"file": "https://fastly.jsdelivr.net/gh/SXP-Simon/profile_assets@main/acgti/characters/GGGG-A.png"
|
||||
},
|
||||
{
|
||||
"code": "ANON",
|
||||
"name": "千早爱音",
|
||||
"mbti": "ESFJ",
|
||||
"file": "https://fastly.jsdelivr.net/gh/SXP-Simon/profile_assets@main/acgti/characters/ANON.png"
|
||||
},
|
||||
{
|
||||
"code": "BAFE",
|
||||
"name": "要乐奈",
|
||||
"mbti": "ISTP",
|
||||
"file": "https://fastly.jsdelivr.net/gh/SXP-Simon/profile_assets@main/acgti/characters/BAFE.png"
|
||||
},
|
||||
{
|
||||
"code": "SOYO",
|
||||
"name": "长崎爽世",
|
||||
"mbti": "ISFJ",
|
||||
"file": "https://fastly.jsdelivr.net/gh/SXP-Simon/profile_assets@main/acgti/characters/SOYO.png"
|
||||
},
|
||||
{
|
||||
"code": "RIKI",
|
||||
"name": "椎名立希",
|
||||
"mbti": "ESTJ",
|
||||
"file": "https://fastly.jsdelivr.net/gh/SXP-Simon/profile_assets@main/acgti/characters/RIKI.png"
|
||||
},
|
||||
{
|
||||
"code": "SAKI",
|
||||
"name": "丰川祥子",
|
||||
"mbti": "ENTJ",
|
||||
"file": "https://fastly.jsdelivr.net/gh/SXP-Simon/profile_assets@main/acgti/characters/SAKI.png"
|
||||
},
|
||||
{
|
||||
"code": "MRTS",
|
||||
"name": "若叶睦",
|
||||
"mbti": "ISTJ",
|
||||
"file": "https://fastly.jsdelivr.net/gh/SXP-Simon/profile_assets@main/acgti/characters/MRTS.png"
|
||||
},
|
||||
{
|
||||
"code": "MRTS-X",
|
||||
"name": "Mortis",
|
||||
"mbti": "INTJ",
|
||||
"file": "https://fastly.jsdelivr.net/gh/SXP-Simon/profile_assets@main/acgti/characters/MRTS-X.png"
|
||||
},
|
||||
{
|
||||
"code": "DLRS",
|
||||
"name": "三角初华 / Doloris",
|
||||
"mbti": "INFJ",
|
||||
"file": "https://fastly.jsdelivr.net/gh/SXP-Simon/profile_assets@main/acgti/characters/DLRS.png"
|
||||
},
|
||||
{
|
||||
"code": "TMRS",
|
||||
"name": "八幡海铃 / Timoris",
|
||||
"mbti": "ISTJ",
|
||||
"file": "https://fastly.jsdelivr.net/gh/SXP-Simon/profile_assets@main/acgti/characters/TMRS.png"
|
||||
},
|
||||
{
|
||||
"code": "AMRS",
|
||||
"name": "祐天寺若麦 / Amoris",
|
||||
"mbti": "ENTP",
|
||||
"file": "https://fastly.jsdelivr.net/gh/SXP-Simon/profile_assets@main/acgti/characters/AMRS.png"
|
||||
},
|
||||
{
|
||||
"code": "DINA",
|
||||
"name": "嘉然",
|
||||
"mbti": "ENFP",
|
||||
"file": "https://fastly.jsdelivr.net/gh/SXP-Simon/profile_assets@main/acgti/characters/DINA.png"
|
||||
},
|
||||
{
|
||||
"code": "TAFI",
|
||||
"name": "永雏塔菲",
|
||||
"mbti": "ENTP",
|
||||
"file": "https://fastly.jsdelivr.net/gh/SXP-Simon/profile_assets@main/acgti/characters/TAFI.png"
|
||||
},
|
||||
{
|
||||
"code": "NEUR",
|
||||
"name": "Neuro-sama",
|
||||
"mbti": "ENTP",
|
||||
"file": "https://fastly.jsdelivr.net/gh/SXP-Simon/profile_assets@main/acgti/characters/NEUR.png"
|
||||
},
|
||||
{
|
||||
"code": "EVIL",
|
||||
"name": "Evil Neuro",
|
||||
"mbti": "INTJ",
|
||||
"file": "https://fastly.jsdelivr.net/gh/SXP-Simon/profile_assets@main/acgti/characters/EVIL.png"
|
||||
},
|
||||
{
|
||||
"code": "RINN",
|
||||
"name": "镜音铃",
|
||||
"mbti": "ESFP",
|
||||
"file": "https://fastly.jsdelivr.net/gh/SXP-Simon/profile_assets@main/acgti/characters/RINN.png"
|
||||
},
|
||||
{
|
||||
"code": "LTYI",
|
||||
"name": "洛天依",
|
||||
"mbti": "ISFP",
|
||||
"file": "https://fastly.jsdelivr.net/gh/SXP-Simon/profile_assets@main/acgti/characters/LTYI.png"
|
||||
},
|
||||
{
|
||||
"code": "YCYO",
|
||||
"name": "月见八千代",
|
||||
"mbti": "ENFJ",
|
||||
"file": "https://fastly.jsdelivr.net/gh/SXP-Simon/profile_assets@main/acgti/characters/YCYO.png"
|
||||
}
|
||||
]
|
||||
}
|
||||
|
After Width: | Height: | Size: 843 KiB |
|
After Width: | Height: | Size: 1.0 MiB |
|
After Width: | Height: | Size: 629 KiB |
|
After Width: | Height: | Size: 673 KiB |
@@ -0,0 +1,36 @@
|
||||
"""NoneBot 版 BotManager:管理 OneBot 适配器实例。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
|
||||
class BotManager:
|
||||
def __init__(self):
|
||||
self._adapters: dict[str, Any] = {}
|
||||
|
||||
def register_adapter(self, adapter: Any, platform_id: str | None = None) -> None:
|
||||
pid = platform_id or getattr(adapter, "platform_id", "onebot") or "onebot"
|
||||
self._adapters[str(pid)] = adapter
|
||||
|
||||
def set_bot_instance(self, bot_instance: Any, platform_id: str | None = None) -> None:
|
||||
"""兼容接口:按平台登记一个适配器。"""
|
||||
pid = platform_id or "onebot"
|
||||
from .adapter import OneBotAdapter
|
||||
self.register_adapter(OneBotAdapter(bot_instance, {"platform_id": pid}), pid)
|
||||
|
||||
def get_adapter(self, platform_id: str | None = None) -> Any | None:
|
||||
if not self._adapters:
|
||||
return None
|
||||
if platform_id and str(platform_id) in self._adapters:
|
||||
return self._adapters[str(platform_id)]
|
||||
return next(iter(self._adapters.values()))
|
||||
|
||||
def get_adapter_platform_id(self, adapter: Any) -> str:
|
||||
return str(getattr(adapter, "platform_id", "onebot"))
|
||||
|
||||
def get_platform_ids(self) -> list[str]:
|
||||
return list(self._adapters.keys())
|
||||
|
||||
def get_platform_count(self) -> int:
|
||||
return len(self._adapters)
|
||||
@@ -0,0 +1,142 @@
|
||||
"""NoneBot 版插件配置加载。
|
||||
|
||||
配置持久化在数据目录 config.json,键值结构与原 AstrBot 配置分组一致,
|
||||
这样可以直接喂给 core 的 ConfigManager(它本质只读嵌套 dict)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
from .core.shared.constants import PLUGIN_NAME
|
||||
from .core.utils.paths import get_data_dir
|
||||
|
||||
|
||||
class ConfigDict(dict):
|
||||
"""兼容原 ConfigManager 的 dict;save_config 会把改动持久化到 config.json。"""
|
||||
|
||||
def save_config(self) -> None:
|
||||
try:
|
||||
config_file().write_text(
|
||||
json.dumps(dict(self), ensure_ascii=False, indent=2, default=str),
|
||||
encoding="utf-8",
|
||||
)
|
||||
except OSError:
|
||||
pass
|
||||
|
||||
|
||||
def _build_from_items(items: dict) -> dict:
|
||||
"""从 schema 的 items 递归生成嵌套默认配置。"""
|
||||
out: dict = {}
|
||||
for key, item in items.items():
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
if item.get("type") == "object" and isinstance(item.get("items"), dict):
|
||||
out[key] = _build_from_items(item["items"])
|
||||
elif "default" in item:
|
||||
out[key] = item["default"]
|
||||
return out
|
||||
|
||||
|
||||
def _schema_defaults() -> dict:
|
||||
"""从插件的 _conf_schema.json 读取完整默认配置。"""
|
||||
schema_path = Path(__file__).resolve().parents[0] / "_conf_schema.json"
|
||||
try:
|
||||
schema = json.loads(schema_path.read_text(encoding="utf-8-sig"))
|
||||
except (OSError, json.JSONDecodeError):
|
||||
return {}
|
||||
defaults: dict = {}
|
||||
for group, spec in schema.items():
|
||||
if isinstance(spec, dict) and isinstance(spec.get("items"), dict):
|
||||
defaults[group] = _build_from_items(spec["items"])
|
||||
return defaults
|
||||
|
||||
|
||||
def _read_config_value(*names: str, default: str = "") -> str:
|
||||
"""依次从 shell 环境变量与 NoneBot driver.config(.env) 读取值。
|
||||
|
||||
NoneBot 的 .env 是通过 pydantic 加载的,不会写入 os.environ,
|
||||
因此必须能从 get_driver().config 的小写属性读取。优先环境变量。
|
||||
"""
|
||||
for n in names:
|
||||
v = os.environ.get(n)
|
||||
if v:
|
||||
return str(v)
|
||||
try:
|
||||
from nonebot import get_driver
|
||||
drv_cfg = get_driver().config
|
||||
except Exception:
|
||||
return default
|
||||
for n in names:
|
||||
attr = n.lower()
|
||||
try:
|
||||
v = getattr(drv_cfg, attr, None)
|
||||
except Exception:
|
||||
v = None
|
||||
if v:
|
||||
return str(v)
|
||||
return default
|
||||
|
||||
|
||||
def _llm_env_values() -> dict:
|
||||
"""用户必填的 LLM 端点/密钥/模型;AstrBot 版是 Provider 体系,插件本身不存这些。"""
|
||||
return {
|
||||
"llm_api_base": _read_config_value("HEXI_LLM_API_BASE"),
|
||||
"llm_api_key": _read_config_value("HEXI_LLM_API_KEY"),
|
||||
"llm_model": _read_config_value("HEXI_LLM_MODEL"),
|
||||
}
|
||||
|
||||
|
||||
def _default_config() -> dict:
|
||||
cfg = _schema_defaults()
|
||||
cfg.setdefault("llm", {}).update(_llm_env_values())
|
||||
cfg.setdefault("basic", {}).setdefault("output_format", ["image"])
|
||||
cfg.setdefault("basic", {}).setdefault("enable_base64_image", True)
|
||||
cfg.setdefault("basic", {}).setdefault("analysis_days", 1)
|
||||
cfg.setdefault("basic", {}).setdefault("max_messages", 1000)
|
||||
cfg.setdefault("basic", {}).setdefault("min_messages_threshold", 50)
|
||||
cfg.setdefault("basic", {}).setdefault("filter_bot_messages", True)
|
||||
cfg.setdefault("analysis_features", {}).setdefault("chat_quality_analysis_enabled", False)
|
||||
cfg.setdefault("incremental", {}).setdefault("incremental_enabled", False)
|
||||
return cfg
|
||||
|
||||
|
||||
def config_file() -> Path:
|
||||
return get_data_dir(PLUGIN_NAME) / "config.json"
|
||||
|
||||
|
||||
def load_config() -> ConfigDict:
|
||||
cfg = ConfigDict(_default_config())
|
||||
f = config_file()
|
||||
if f.exists():
|
||||
try:
|
||||
user = json.loads(f.read_text(encoding="utf-8"))
|
||||
for group, values in user.items():
|
||||
if isinstance(values, dict):
|
||||
cfg.setdefault(group, {}).update(values)
|
||||
else:
|
||||
cfg[group] = values
|
||||
except (OSError, json.JSONDecodeError) as e:
|
||||
_log_warn(f"读取配置文件失败: {e}")
|
||||
# 环境变量/NoneBot .env 永远优先于持久化 config.json,避免被空值覆盖
|
||||
cfg.setdefault("llm", {}).update(_llm_env_values())
|
||||
return cfg
|
||||
|
||||
|
||||
def save_config(cfg: dict) -> None:
|
||||
try:
|
||||
config_file().write_text(
|
||||
json.dumps(cfg, ensure_ascii=False, indent=2), encoding="utf-8"
|
||||
)
|
||||
except OSError as e:
|
||||
_log_warn(f"保存配置文件失败: {e}")
|
||||
|
||||
|
||||
def _log_warn(msg: str) -> None:
|
||||
try:
|
||||
from .core.utils.logger import logger
|
||||
logger.warning(msg)
|
||||
except Exception:
|
||||
pass
|
||||
@@ -0,0 +1,17 @@
|
||||
"""
|
||||
群日常分析插件 - 源代码包
|
||||
|
||||
本包包含插件的核心实现,采用 DDD (领域驱动设计) 架构:
|
||||
- application: 应用层 - 编排领域服务,处理用例
|
||||
- domain: 领域层 - 核心业务逻辑,平台无关
|
||||
- infrastructure: 基础设施层 - 外部服务适配
|
||||
- shared: 共享组件 - 跨层使用的工具和常量
|
||||
|
||||
遗留模块(渐进式迁移中):
|
||||
- analysis: 分析器实现
|
||||
- core: 核心组件
|
||||
- reports: 报告生成
|
||||
- scheduler: 定时任务
|
||||
- utils: 工具函数
|
||||
- visualization: 可视化组件
|
||||
"""
|
||||
@@ -0,0 +1 @@
|
||||
# 应用层 - 编排和用例
|
||||
@@ -0,0 +1,5 @@
|
||||
"""命令相关应用服务。"""
|
||||
|
||||
from .template_command_service import TemplateCommandService
|
||||
|
||||
__all__ = ["TemplateCommandService"]
|
||||
@@ -0,0 +1,117 @@
|
||||
"""模板管理相关命令服务。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
|
||||
from astrbot.api.message_components import Image, Node, Nodes, Plain
|
||||
|
||||
|
||||
class TemplateCommandService:
|
||||
"""封装模板命令的文件系统与消息构建逻辑。"""
|
||||
|
||||
_CIRCLE_NUMBERS = ["①", "②", "③", "④", "⑤", "⑥", "⑦", "⑧", "⑨", "⑩"]
|
||||
|
||||
def __init__(self, plugin_root: str):
|
||||
self.plugin_root = plugin_root
|
||||
|
||||
def resolve_template_base_dir(self) -> str:
|
||||
"""解析报告模板目录(兼容新旧目录结构)。"""
|
||||
candidate_dirs = [
|
||||
os.path.join(
|
||||
self.plugin_root, "src", "infrastructure", "reporting", "templates"
|
||||
),
|
||||
os.path.join(self.plugin_root, "src", "reports", "templates"),
|
||||
]
|
||||
for candidate in candidate_dirs:
|
||||
if os.path.isdir(candidate):
|
||||
return candidate
|
||||
return candidate_dirs[0]
|
||||
|
||||
def resolve_template_preview_path(self, template_name: str) -> str | None:
|
||||
"""解析模板预览图路径。"""
|
||||
candidate_paths = [
|
||||
os.path.join(self.plugin_root, "assets", f"{template_name}-demo.jpg"),
|
||||
]
|
||||
for candidate in candidate_paths:
|
||||
if os.path.exists(candidate):
|
||||
return candidate
|
||||
return None
|
||||
|
||||
async def list_available_templates(self) -> list[str]:
|
||||
"""列出所有可用模板。"""
|
||||
template_base_dir = self.resolve_template_base_dir()
|
||||
|
||||
def _list_templates_sync() -> list[str]:
|
||||
if os.path.exists(template_base_dir):
|
||||
return sorted(
|
||||
[
|
||||
d
|
||||
for d in os.listdir(template_base_dir)
|
||||
if os.path.isdir(os.path.join(template_base_dir, d))
|
||||
and not d.startswith("__")
|
||||
]
|
||||
)
|
||||
return []
|
||||
|
||||
return await asyncio.to_thread(_list_templates_sync)
|
||||
|
||||
async def template_exists(self, template_name: str) -> bool:
|
||||
"""检查模板目录是否存在。"""
|
||||
template_dir = os.path.join(self.resolve_template_base_dir(), template_name)
|
||||
return await asyncio.to_thread(os.path.exists, template_dir)
|
||||
|
||||
def parse_template_input(
|
||||
self, template_input: str, available_templates: list[str]
|
||||
) -> tuple[str | None, str | None]:
|
||||
"""解析模板输入(支持模板名或序号)。"""
|
||||
if not template_input:
|
||||
return None, "❌ 模板参数不能为空"
|
||||
|
||||
if template_input.isdigit():
|
||||
index = int(template_input)
|
||||
if 1 <= index <= len(available_templates):
|
||||
return available_templates[index - 1], None
|
||||
return (
|
||||
None,
|
||||
f"❌ 无效的序号 '{template_input}',有效范围: 1-{len(available_templates)}",
|
||||
)
|
||||
|
||||
return template_input, None
|
||||
|
||||
def build_template_preview_nodes(
|
||||
self,
|
||||
available_templates: list[str],
|
||||
current_template: str,
|
||||
bot_id: str,
|
||||
) -> Nodes:
|
||||
"""构建模板预览的合并消息节点。"""
|
||||
node_list = []
|
||||
|
||||
header_content = [
|
||||
Plain(
|
||||
f"🎨 可用报告模板列表\n📌 当前使用: {current_template}\n💡 使用 /设置模板 [序号] 切换"
|
||||
)
|
||||
]
|
||||
node_list.append(Node(uin=bot_id, name="模板预览", content=header_content))
|
||||
|
||||
for index, template_name in enumerate(available_templates):
|
||||
current_mark = " ✅" if template_name == current_template else ""
|
||||
num_label = (
|
||||
self._CIRCLE_NUMBERS[index]
|
||||
if index < len(self._CIRCLE_NUMBERS)
|
||||
else f"({index + 1})"
|
||||
)
|
||||
|
||||
node_content = [Plain(f"{num_label} {template_name}{current_mark}")]
|
||||
preview_image_path = self.resolve_template_preview_path(template_name)
|
||||
if preview_image_path:
|
||||
node_content.append(Image.fromFileSystem(preview_image_path))
|
||||
else:
|
||||
cdn_url = f"https://fastly.jsdelivr.net/gh/SXP-Simon/astrbot_plugin_qq_group_daily_analysis@main/assets/{template_name}-demo.jpg"
|
||||
node_content.append(Image.fromURL(cdn_url))
|
||||
|
||||
node_list.append(Node(uin=bot_id, name=template_name, content=node_content))
|
||||
|
||||
return Nodes(node_list)
|
||||
@@ -0,0 +1,617 @@
|
||||
import mimetypes
|
||||
from contextlib import nullcontext
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from astrbot.api.star import Context
|
||||
|
||||
from ...infrastructure.analysis.llm_analyzer import LLMAnalyzer
|
||||
from ...infrastructure.config.config_manager import ConfigManager
|
||||
from ...infrastructure.drawing.drawing_client import (
|
||||
DrawingClient,
|
||||
ImageDownloadFailedError,
|
||||
)
|
||||
from ...shared.trace_context import TraceContext
|
||||
from ...utils.logger import logger
|
||||
|
||||
|
||||
class ComicApplicationService:
|
||||
"""
|
||||
负责统筹每日群漫画的生成流程:
|
||||
1. 调用 LLMAnalyzer 将群聊话题生成拼贴分镜提示词。
|
||||
2. 调用 DrawingClient 直接生成单张连环漫画长图。
|
||||
3. 返回图片数据供外部上传。
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
llm_analyzer: LLMAnalyzer,
|
||||
drawing_client: DrawingClient,
|
||||
config_manager: ConfigManager,
|
||||
plugin_data_dir: Path,
|
||||
context: Context | None = None,
|
||||
):
|
||||
self.llm_analyzer = llm_analyzer
|
||||
self.drawing_client = drawing_client
|
||||
self.config_manager = config_manager
|
||||
self.plugin_data_dir = plugin_data_dir
|
||||
self.context = context
|
||||
|
||||
async def generate_comic(
|
||||
self,
|
||||
topics: list[dict],
|
||||
group_id: str,
|
||||
umo: str | None = None,
|
||||
) -> tuple[bytes | None, str | None]:
|
||||
"""
|
||||
生成漫画并返回图片字节数据。
|
||||
|
||||
Returns:
|
||||
(comic_bytes, fallback_url):
|
||||
- comic_bytes: 生成成功时为图片字节,失败时为 None。
|
||||
- fallback_url: 图片 API 返回了 URL 但下载失败时为该 URL,其他情况为 None。
|
||||
"""
|
||||
if not self.config_manager.get_enable_daily_comic():
|
||||
return None, None
|
||||
|
||||
character = self.config_manager.get_selected_comic_character()
|
||||
character_name = (
|
||||
str(character.get("name", "")).strip() if character else ""
|
||||
) or "默认配置"
|
||||
persona_id = self.config_manager.get_comic_character_persona_id(character)
|
||||
prompt_template = self.config_manager.get_comic_character_storyboard_prompt(
|
||||
character
|
||||
)
|
||||
logger.info(
|
||||
f"[Comic] 开始为群 {group_id} 生成每日漫画,角色方案: {character_name}"
|
||||
)
|
||||
|
||||
trace = TraceContext.current()
|
||||
|
||||
# 1. 提取分镜和金句
|
||||
sb_ctx = trace.span("COMIC_STORYBOARD") if trace else nullcontext()
|
||||
with sb_ctx as sb_rec:
|
||||
(
|
||||
storyboards,
|
||||
storyboard_usage,
|
||||
) = await self.llm_analyzer.analyze_comic_storyboards(
|
||||
topics,
|
||||
umo,
|
||||
persona_id=persona_id or None,
|
||||
prompt_template=prompt_template or None,
|
||||
)
|
||||
if sb_rec and isinstance(sb_rec, dict):
|
||||
sb_prompts: dict[str, Any] = {}
|
||||
if trace and trace.metadata.get("llm_prompts"):
|
||||
for k, p in trace.metadata["llm_prompts"].items():
|
||||
if "comic" in k or k == "comic_storyboards":
|
||||
sb_prompts[k] = p
|
||||
if not sb_prompts and storyboards:
|
||||
sb_prompts["comic_storyboards"] = {
|
||||
"prompt": prompt_template
|
||||
or "自动从群聊话题中提取漫画多格分镜与生图提示词",
|
||||
"system_prompt": f"漫画分镜师 | 人格: {persona_id or character_name}",
|
||||
"completion": storyboards[0].get("scene", "")
|
||||
if storyboards
|
||||
else "",
|
||||
"prompt_tokens": getattr(storyboard_usage, "prompt_tokens", 0),
|
||||
"completion_tokens": getattr(
|
||||
storyboard_usage, "completion_tokens", 0
|
||||
),
|
||||
"tokens": getattr(storyboard_usage, "total_tokens", 0),
|
||||
}
|
||||
sb_rec.setdefault("payload", {}).update(
|
||||
{
|
||||
"character_name": character_name,
|
||||
"topics_count": len(topics),
|
||||
"storyboards_count": len(storyboards) if storyboards else 0,
|
||||
"prompt_tokens": getattr(storyboard_usage, "prompt_tokens", 0),
|
||||
"completion_tokens": getattr(
|
||||
storyboard_usage, "completion_tokens", 0
|
||||
),
|
||||
"total_tokens": getattr(storyboard_usage, "total_tokens", 0),
|
||||
"prompts": sb_prompts,
|
||||
}
|
||||
)
|
||||
if not storyboards:
|
||||
sb_rec["payload"]["warning"] = "未能从群聊话题中提取出漫画分镜"
|
||||
|
||||
if not storyboards:
|
||||
logger.warning(
|
||||
f"[Comic] 群 {group_id} 未能提取到任何金句分镜,取消漫画生成。"
|
||||
)
|
||||
return None, None
|
||||
|
||||
logger.info("[Comic] 成功提取到全景分镜提示词,开始调用绘画 API...")
|
||||
|
||||
# 2. 直接生成一张图片
|
||||
scene_prompt = storyboards[0].get("scene", "")
|
||||
if not scene_prompt:
|
||||
logger.error("[Comic] 提取到的场景提示词为空,取消漫画生成。")
|
||||
return None, None
|
||||
|
||||
logger.debug(f"[Comic] 漫画 Prompt 已生成,长度: {len(scene_prompt)}")
|
||||
|
||||
# 3. 加载当前角色方案配置的全部参考图。
|
||||
images_data = []
|
||||
reference_image_paths = self.config_manager.get_drawing_reference_images()
|
||||
for reference_image_path in reference_image_paths:
|
||||
reference_image = await self._fetch_reference_image(reference_image_path)
|
||||
if reference_image:
|
||||
images_data.append(reference_image)
|
||||
logger.info(f"[Comic] 已加载参考图: {Path(reference_image_path).name}")
|
||||
else:
|
||||
logger.warning(
|
||||
f"[Comic] 无法加载参考图: {Path(reference_image_path).name}"
|
||||
)
|
||||
|
||||
draw_ctx = trace.span("COMIC_DRAWING") if trace else nullcontext()
|
||||
with draw_ctx as draw_rec:
|
||||
# 4. 若配置为外部绘图后端,优先走对应插件出图
|
||||
backend = self.config_manager.get_drawing_backend()
|
||||
if backend in {"general_plugin", "big_banana"}:
|
||||
if backend == "general_plugin":
|
||||
external_comic_bytes = await self._generate_via_general_plugin(
|
||||
scene_prompt, images_data
|
||||
)
|
||||
else:
|
||||
external_comic_bytes = await self._generate_via_big_banana(
|
||||
scene_prompt, images_data
|
||||
)
|
||||
if external_comic_bytes and not any(
|
||||
external_comic_bytes == reference[0] for reference in images_data
|
||||
):
|
||||
logger.info(
|
||||
f"[Comic] 漫画生成成功({backend} 后端),大小: {len(external_comic_bytes)} bytes"
|
||||
)
|
||||
if draw_rec and isinstance(draw_rec, dict):
|
||||
draw_rec.setdefault("payload", {}).update(
|
||||
{
|
||||
"backend": backend,
|
||||
"scene_prompt_len": len(scene_prompt),
|
||||
"reference_images_count": len(images_data),
|
||||
"image_bytes": len(external_comic_bytes),
|
||||
"success": True,
|
||||
"prompts": {
|
||||
"comic_drawing": {
|
||||
"prompt": scene_prompt,
|
||||
"system_prompt": f"绘图后端: {backend} | 角色方案: {character_name} | 参考图数: {len(images_data)}",
|
||||
"completion": f"出图完成(体积: {round(len(external_comic_bytes) / 1024, 1)} KB)",
|
||||
"provider_type": backend,
|
||||
}
|
||||
},
|
||||
}
|
||||
)
|
||||
return external_comic_bytes, None
|
||||
if external_comic_bytes:
|
||||
logger.warning(
|
||||
f"[Comic] {backend} 后端原样返回了参考图,拒绝发送并回退内置绘图后端。"
|
||||
)
|
||||
if not self.config_manager.get_drawing_external_fallback():
|
||||
logger.warning(
|
||||
f"[Comic] {backend} 后端未产出结果,且已禁用回退内置后端,取消漫画生成。"
|
||||
)
|
||||
if draw_rec and isinstance(draw_rec, dict):
|
||||
draw_rec.setdefault("payload", {}).update(
|
||||
{
|
||||
"backend": backend,
|
||||
"error": f"{backend} 后端未产出结果且禁用回退",
|
||||
"success": False,
|
||||
}
|
||||
)
|
||||
return None, None
|
||||
logger.warning(f"[Comic] {backend} 后端未产出结果,回退内置绘图后端。")
|
||||
|
||||
# 5. 内置绘图后端未配置时直接取消,避免空跑
|
||||
if not self.config_manager.get_drawing_provider_configs():
|
||||
logger.warning(
|
||||
"[Comic] 未配置绘图供应商(drawing_provider_overrides),取消漫画生成。"
|
||||
)
|
||||
if draw_rec and isinstance(draw_rec, dict):
|
||||
draw_rec.setdefault("payload", {}).update(
|
||||
{
|
||||
"backend": "builtin",
|
||||
"error": "未配置绘图供应商",
|
||||
"success": False,
|
||||
}
|
||||
)
|
||||
return None, None
|
||||
|
||||
# 6. 调用绘图 API,捕获"有 URL 但下载失败"的情况
|
||||
fallback_url: str | None = None
|
||||
try:
|
||||
(
|
||||
final_comic_bytes,
|
||||
last_error,
|
||||
) = await self.drawing_client.generate_image(
|
||||
scene_prompt, images_data=images_data or None
|
||||
)
|
||||
except ImageDownloadFailedError as exc:
|
||||
logger.warning(
|
||||
f"[Comic] 图片下载失败,保留 fallback URL: {exc.fallback_url}"
|
||||
)
|
||||
if draw_rec and isinstance(draw_rec, dict):
|
||||
draw_rec.setdefault("payload", {}).update(
|
||||
{
|
||||
"backend": "builtin",
|
||||
"fallback_url": exc.fallback_url,
|
||||
"error": "图片下载失败,使用 fallback URL 发送",
|
||||
}
|
||||
)
|
||||
return None, exc.fallback_url
|
||||
|
||||
if final_comic_bytes and any(
|
||||
final_comic_bytes == reference[0] for reference in images_data
|
||||
):
|
||||
logger.warning("[Comic] 内建绘图原样返回了参考图,拒绝发送。")
|
||||
final_comic_bytes = None
|
||||
last_error = "绘图服务原样返回了参考图"
|
||||
|
||||
exception_keywords = (
|
||||
self.config_manager.get_drawing_output_exception_retry_keywords()
|
||||
)
|
||||
should_rewrite_prompt = bool(
|
||||
last_error
|
||||
and any(
|
||||
keyword in last_error for keyword in exception_keywords if keyword
|
||||
)
|
||||
)
|
||||
if not final_comic_bytes and last_error and should_rewrite_prompt:
|
||||
logger.info(
|
||||
f"[Comic] 画图重试已用尽,请求 LLM 分析报错并重写 Prompt: {last_error}"
|
||||
)
|
||||
new_prompt = await self.llm_analyzer.analyze_retry_prompt(
|
||||
scene_prompt, last_error, umo
|
||||
)
|
||||
if new_prompt:
|
||||
logger.info("[Comic] 获取到重写后的 Prompt,进行最后一次尝试...")
|
||||
try:
|
||||
final_comic_bytes, _ = await self.drawing_client.generate_image(
|
||||
new_prompt,
|
||||
images_data=images_data or None,
|
||||
disable_retry=True,
|
||||
)
|
||||
except ImageDownloadFailedError as exc:
|
||||
logger.warning(
|
||||
f"[Comic] 重写 Prompt 后图片下载仍失败,保留 fallback URL: {exc.fallback_url}"
|
||||
)
|
||||
if draw_rec and isinstance(draw_rec, dict):
|
||||
draw_rec.setdefault("payload", {}).update(
|
||||
{
|
||||
"backend": "builtin",
|
||||
"fallback_url": exc.fallback_url,
|
||||
"error": "重写 Prompt 后下载仍失败",
|
||||
}
|
||||
)
|
||||
return None, exc.fallback_url
|
||||
if final_comic_bytes and any(
|
||||
final_comic_bytes == reference[0] for reference in images_data
|
||||
):
|
||||
logger.warning(
|
||||
"[Comic] 重写 Prompt 后仍原样返回参考图,拒绝发送。"
|
||||
)
|
||||
final_comic_bytes = None
|
||||
|
||||
if draw_rec and isinstance(draw_rec, dict):
|
||||
draw_prompts = {
|
||||
"comic_drawing": {
|
||||
"prompt": scene_prompt,
|
||||
"system_prompt": f"绘图引擎: {backend} | 角色方案: {character_name} | 参考图数: {len(images_data)}",
|
||||
"completion": f"出图完成(体积: {round(len(final_comic_bytes) / 1024, 1)} KB)"
|
||||
if final_comic_bytes
|
||||
else (
|
||||
f"出图未产出: {last_error}" if last_error else "未产出图像"
|
||||
),
|
||||
"provider_type": backend,
|
||||
}
|
||||
}
|
||||
draw_rec.setdefault("payload", {}).update(
|
||||
{
|
||||
"backend": backend,
|
||||
"scene_prompt_len": len(scene_prompt),
|
||||
"reference_images_count": len(images_data),
|
||||
"image_bytes": len(final_comic_bytes)
|
||||
if final_comic_bytes
|
||||
else 0,
|
||||
"success": bool(final_comic_bytes),
|
||||
"last_error": last_error,
|
||||
"prompts": draw_prompts,
|
||||
}
|
||||
)
|
||||
|
||||
if final_comic_bytes:
|
||||
logger.info(
|
||||
f"[Comic] 漫画生成成功,大小: {len(final_comic_bytes)} bytes"
|
||||
)
|
||||
else:
|
||||
logger.error("[Comic] 漫画生成最终失败。")
|
||||
|
||||
return final_comic_bytes, fallback_url
|
||||
|
||||
async def _generate_via_general_plugin(
|
||||
self,
|
||||
scene_prompt: str,
|
||||
images_data: list[tuple[bytes, str]] | None,
|
||||
) -> bytes | None:
|
||||
"""通过「通用生图」插件的公共 API 生成漫画。
|
||||
|
||||
未安装、未激活、未配置 API 或调用失败时返回 None,由调用方回退内置 DrawingClient。
|
||||
|
||||
Returns:
|
||||
生成图片的二进制数据;失败时返回 None。
|
||||
"""
|
||||
if self.context is None:
|
||||
logger.debug("[Comic] 未注入插件 Context,跳过通用生图后端。")
|
||||
return None
|
||||
try:
|
||||
meta = self.context.get_registered_star("astrbot_plugin_image_generation")
|
||||
except Exception as exc:
|
||||
logger.debug(f"[Comic] 获取通用生图插件注册信息失败: {exc}")
|
||||
return None
|
||||
image_plugin = meta.star_cls if meta and meta.activated else None
|
||||
if image_plugin is None:
|
||||
logger.warning(
|
||||
"[Comic] 未检测到已激活的「通用生图」插件,回退内置绘图后端。"
|
||||
)
|
||||
return None
|
||||
|
||||
public_api = getattr(image_plugin, "public_api", None)
|
||||
if public_api is None:
|
||||
logger.warning("[Comic] 通用生图插件未暴露 public_api,回退内置绘图后端。")
|
||||
return None
|
||||
|
||||
try:
|
||||
logger.info("[Comic] 通过「通用生图」插件公共 API 生成漫画...")
|
||||
result = await public_api.generate_image_files(
|
||||
prompt=scene_prompt,
|
||||
source="群分析插件",
|
||||
aspect_ratio="16:9",
|
||||
reference_image_data=images_data,
|
||||
timeout_seconds=600,
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.error(f"[Comic] 通用生图后端调用异常: {exc}")
|
||||
return None
|
||||
|
||||
if not getattr(result, "ok", False):
|
||||
code = str(getattr(result, "code", ""))
|
||||
message = getattr(result, "message", "") or getattr(result, "error", "")
|
||||
hint = ""
|
||||
if code == "prompt_blocked":
|
||||
hint = "(提示词被通用生图插件安全审核拦截,可调整其审核配置或精简 scene 提示词)"
|
||||
elif code == "api_key_missing":
|
||||
hint = "(通用生图插件未配置 API Key,需先在通用生图插件中配置)"
|
||||
elif code == "timeout":
|
||||
hint = "(等待通用生图任务结果超时)"
|
||||
elif code == "rate_limited":
|
||||
hint = "(命中通用生图插件额度/频率限制)"
|
||||
logger.warning(f"[Comic] 通用生图后端失败 [{code}]: {message}{hint}")
|
||||
return None
|
||||
|
||||
paths = list(getattr(result, "paths", None) or [])
|
||||
if not paths:
|
||||
logger.warning(
|
||||
"[Comic] 通用生图后端未返回图片路径(可能参考图被忽略或结果为空,请检查通用生图插件配置与参考图大小限制)。"
|
||||
)
|
||||
return None
|
||||
try:
|
||||
return Path(paths[0]).read_bytes()
|
||||
except OSError as exc:
|
||||
logger.warning(f"[Comic] 读取通用生图后端结果失败: {exc}")
|
||||
return None
|
||||
|
||||
async def _generate_via_big_banana(
|
||||
self,
|
||||
scene_prompt: str,
|
||||
images_data: list[tuple[bytes, str]] | None,
|
||||
) -> bytes | None:
|
||||
"""通过「大香蕉」插件的绘图管线生成漫画。
|
||||
|
||||
大香蕉支持 Gemini、SiliconFlow、OpenAI 等多家提供商。
|
||||
|
||||
Returns:
|
||||
生成图片的二进制数据;失败时返回 None。
|
||||
"""
|
||||
if self.context is None:
|
||||
logger.debug("[Comic] 未注入插件 Context,跳过「大香蕉」后端。")
|
||||
return None
|
||||
try:
|
||||
meta = self.context.get_registered_star("astrbot_plugin_big_banana")
|
||||
except Exception as exc:
|
||||
logger.debug(f"[Comic] 获取「大香蕉」插件注册信息失败: {exc}")
|
||||
return None
|
||||
plugin = meta.star_cls if meta and meta.activated else None
|
||||
if plugin is None:
|
||||
logger.warning("[Comic] 未检测到已激活的「大香蕉」插件,回退内置绘图后端。")
|
||||
return None
|
||||
|
||||
drawing_pipeline = getattr(plugin, "drawing_pipeline", None)
|
||||
if drawing_pipeline is None:
|
||||
logger.warning("[Comic] 「大香蕉」插件未初始化绘图管线,回退内置绘图后端。")
|
||||
return None
|
||||
|
||||
ImageResource = self._import_big_banana_image_resource(plugin)
|
||||
if ImageResource is None:
|
||||
logger.warning("[Comic] 无法导入「大香蕉」图片资源类型,回退内置绘图后端。")
|
||||
return None
|
||||
|
||||
image_list = None
|
||||
if images_data:
|
||||
image_list = []
|
||||
for img_bytes, _mime in images_data:
|
||||
resource = ImageResource.from_bytes(img_bytes)
|
||||
if resource:
|
||||
image_list.append(resource)
|
||||
if not image_list:
|
||||
logger.warning(
|
||||
"[Comic] 参考图无法解析为「大香蕉」图片资源,将不带参考图生成。"
|
||||
)
|
||||
|
||||
params: dict[str, Any] = {
|
||||
"prompt": scene_prompt,
|
||||
"capability": "image_generation",
|
||||
"sub_brain": False,
|
||||
"url": False,
|
||||
"aspect_ratio": "16:9",
|
||||
"image_size": "1K",
|
||||
}
|
||||
|
||||
try:
|
||||
logger.info("[Comic] 通过「大香蕉」插件绘图管线生成漫画...")
|
||||
result = await drawing_pipeline.run(params, image_list)
|
||||
except Exception as exc:
|
||||
logger.error(f"[Comic] 「大香蕉」绘图管线调用异常: {exc}")
|
||||
return None
|
||||
|
||||
if getattr(result, "error_message", None):
|
||||
logger.warning(f"[Comic] 「大香蕉」生成失败: {result.error_message}")
|
||||
return None
|
||||
|
||||
images = getattr(result, "images", None) or []
|
||||
if not images:
|
||||
logger.warning("[Comic] 「大香蕉」未返回图片。")
|
||||
return None
|
||||
try:
|
||||
image_bytes = images[0].bytes
|
||||
except Exception as exc:
|
||||
logger.warning(f"[Comic] 读取「大香蕉」生成结果失败: {exc}")
|
||||
return None
|
||||
if not image_bytes:
|
||||
logger.warning("[Comic] 「大香蕉」返回的图片为空。")
|
||||
return None
|
||||
return image_bytes
|
||||
|
||||
@staticmethod
|
||||
def _import_big_banana_image_resource(plugin: Any):
|
||||
"""导入「大香蕉」插件的 ImageResource 类型。
|
||||
|
||||
AstrBot 以 ``data.plugins.<插件名>.main`` 形式加载插件,模块名并非
|
||||
``astrbot_plugin_big_banana``,因此先从插件类的模块路径推导包名导入;
|
||||
推导失败时回退直接导入,兼容 pip 安装或测试环境注入的场景。
|
||||
|
||||
Args:
|
||||
plugin: 已激活的大香蕉插件实例。
|
||||
|
||||
Returns:
|
||||
ImageResource 类型;无法导入时返回 None。
|
||||
"""
|
||||
import importlib
|
||||
|
||||
module_name = getattr(type(plugin), "__module__", "") or ""
|
||||
candidate_modules = []
|
||||
if module_name and "." in module_name:
|
||||
candidate_modules.append(module_name.rsplit(".", 1)[0] + ".core.schemas")
|
||||
candidate_modules.append("astrbot_plugin_big_banana.core.schemas")
|
||||
|
||||
for module_path in candidate_modules:
|
||||
try:
|
||||
schemas_module = importlib.import_module(module_path)
|
||||
except Exception:
|
||||
continue
|
||||
image_resource = getattr(schemas_module, "ImageResource", None)
|
||||
if image_resource is not None:
|
||||
return image_resource
|
||||
return None
|
||||
|
||||
async def _fetch_reference_image(self, image_ref: str) -> tuple[bytes, str] | None:
|
||||
"""从插件目录、AstrBot files、本地路径、HTTP URL 或 Base64 Data URL 获取已选参考图。
|
||||
|
||||
Args:
|
||||
image_ref: 包含文件路径、URL 或 Base64 Data URL 的字符串。
|
||||
|
||||
Returns:
|
||||
图片字节和 MIME 类型;加载失败时返回 None。
|
||||
"""
|
||||
import base64
|
||||
|
||||
if not image_ref or not isinstance(image_ref, str):
|
||||
return None
|
||||
|
||||
image_ref = image_ref.strip()
|
||||
|
||||
# 1. 支持 Data URL 格式 (data:image/png;base64,xxxx)
|
||||
if image_ref.startswith("data:image/"):
|
||||
try:
|
||||
header, b64_data = image_ref.split(",", 1)
|
||||
mime_type = header.split(";")[0].replace("data:", "").strip()
|
||||
return base64.b64decode(b64_data), mime_type or "image/png"
|
||||
except Exception as e:
|
||||
logger.error(f"[Comic] 解析 Data URL 参考图失败: {e}")
|
||||
return None
|
||||
|
||||
# 2. 支持 base64:// 格式
|
||||
if image_ref.startswith("base64://"):
|
||||
try:
|
||||
return base64.b64decode(image_ref[9:]), "image/png"
|
||||
except Exception as e:
|
||||
logger.error(f"[Comic] 解析 Base64 参考图失败: {e}")
|
||||
return None
|
||||
|
||||
# 3. 支持 HTTP / HTTPS 远程 URL
|
||||
if image_ref.startswith(("http://", "https://")):
|
||||
try:
|
||||
import httpx
|
||||
|
||||
async with httpx.AsyncClient(timeout=15.0) as client:
|
||||
resp = await client.get(image_ref)
|
||||
if resp.status_code == 200:
|
||||
mime_type = resp.headers.get("content-type", "image/png").split(
|
||||
";"
|
||||
)[0]
|
||||
return resp.content, mime_type
|
||||
except Exception as e:
|
||||
logger.error(f"[Comic] 下载远程参考图失败 {image_ref}: {e}")
|
||||
return None
|
||||
|
||||
# 4. 支持本地文件路径(插件数据目录、插件代码目录、AstrBot 根目录、或绝对路径)
|
||||
try:
|
||||
clean_rel = image_ref.lstrip("/\\")
|
||||
candidate_paths: list[Path] = []
|
||||
if hasattr(self, "plugin_data_dir") and self.plugin_data_dir:
|
||||
p_data = Path(self.plugin_data_dir)
|
||||
candidate_paths.extend(
|
||||
[
|
||||
p_data / clean_rel,
|
||||
p_data / "reference_images" / clean_rel,
|
||||
]
|
||||
)
|
||||
|
||||
candidate_paths.append(
|
||||
Path.cwd()
|
||||
/ "data"
|
||||
/ "plugin_data"
|
||||
/ "astrbot_plugin_qq_group_daily_analysis"
|
||||
/ clean_rel
|
||||
)
|
||||
candidate_paths.append(Path(clean_rel))
|
||||
|
||||
for p in candidate_paths:
|
||||
if p.is_file():
|
||||
guessed_type, _ = mimetypes.guess_type(p.name)
|
||||
return p.read_bytes(), guessed_type or "image/png"
|
||||
|
||||
# 模糊查找纯文件名
|
||||
filename = Path(clean_rel).name
|
||||
search_roots: list[Path] = []
|
||||
if hasattr(self, "plugin_data_dir") and self.plugin_data_dir:
|
||||
search_roots.append(Path(self.plugin_data_dir))
|
||||
search_roots.append(
|
||||
Path.cwd()
|
||||
/ "data"
|
||||
/ "plugin_data"
|
||||
/ "astrbot_plugin_qq_group_daily_analysis"
|
||||
)
|
||||
|
||||
for root in search_roots:
|
||||
files_dir = root / "files"
|
||||
if files_dir.exists():
|
||||
for match in files_dir.rglob(filename):
|
||||
if match.is_file():
|
||||
guessed_type, _ = mimetypes.guess_type(match.name)
|
||||
return match.read_bytes(), guessed_type or "image/png"
|
||||
|
||||
logger.warning(f"[Comic] 找不到已选参考图: {image_ref}")
|
||||
return None
|
||||
except Exception as exc:
|
||||
logger.error(f"[Comic] 获取已选参考图失败 {image_ref}: {exc}")
|
||||
return None
|
||||
@@ -0,0 +1,515 @@
|
||||
import re
|
||||
from collections import Counter, OrderedDict
|
||||
|
||||
from astrbot.api.event import AstrMessageEvent
|
||||
from astrbot.api.star import Context
|
||||
|
||||
from ...infrastructure.persistence.platform_group_registry import PlatformGroupRegistry
|
||||
from ...utils.logger import logger
|
||||
|
||||
_QQ_OFFICIAL_PLATFORM_NAMES = frozenset({"qq_official", "qq_official_webhook"})
|
||||
_QQ_OFFICIAL_MENTION_PATTERN = re.compile(r"<@!?([A-Za-z0-9_-]+)>")
|
||||
_LOCAL_HISTORY_MAX_MESSAGES = 10000
|
||||
|
||||
|
||||
class MessageProcessingService:
|
||||
"""
|
||||
消息处理服务
|
||||
|
||||
解析收到的群消息事件,提取内容与发送者信息,持久化历史记录,
|
||||
并维护事件驱动平台(Telegram、QQ 官方等)的群组注册表。
|
||||
QQ 官方平台特有的重复消息去重逻辑也在本服务中处理。
|
||||
|
||||
职责:
|
||||
1. 解析消息内容(文本、图片、@提及等)
|
||||
2. 解析发送者展示名(跨平台兼容)
|
||||
3. 存储消息历史
|
||||
4. 维护群组注册表,供调度器做群组发现(Telegram、QQ 官方等事件驱动平台)
|
||||
5. QQ 官方事件消息去重(按 message_id 预占 + 确认机制)
|
||||
"""
|
||||
|
||||
def __init__(self, context: Context, group_registry: PlatformGroupRegistry):
|
||||
self.context = context
|
||||
self.group_registry = group_registry
|
||||
self._seen_event_ids: OrderedDict[str, None] = OrderedDict()
|
||||
self._inflight_event_ids: set[str] = set()
|
||||
self._seen_event_ids_limit = 4096
|
||||
# AstrBot 4.26.x 的消息历史接口尚未提供 max_messages 参数。
|
||||
# 首次探测到旧签名后缓存结果,避免每条消息都触发一次失败调用。
|
||||
self._supports_history_max_messages: bool | None = None
|
||||
|
||||
async def process_message(self, event: AstrMessageEvent) -> bool:
|
||||
"""
|
||||
处理并在历史记录中存储消息。
|
||||
被 main.py 的 Telegram 和 QQ 官方消息拦截器共同调用。
|
||||
|
||||
Args:
|
||||
event: AstrBot 消息事件
|
||||
|
||||
Raises:
|
||||
ValueError: 当必要数据无法获取时
|
||||
RuntimeError: 当消息内容为空时
|
||||
"""
|
||||
# 1. 获取群组 ID(必需)
|
||||
group_id = self._get_group_id_from_event(event)
|
||||
if not group_id:
|
||||
raise ValueError("无法获取群组 ID,拒绝存储消息")
|
||||
|
||||
# 2. 获取发送者 ID(必需)
|
||||
sender_id = event.get_sender_id()
|
||||
if not sender_id:
|
||||
raise ValueError(f"群 {group_id}: 无法获取发送者 ID,拒绝存储消息")
|
||||
sender_id = str(sender_id)
|
||||
|
||||
# 3. 获取发送者名称(昵称优先,必要时回退)
|
||||
sender_name = self._resolve_sender_name(event, sender_id)
|
||||
|
||||
# 4. 获取平台 ID(必需)
|
||||
platform_id = event.get_platform_id()
|
||||
if not platform_id:
|
||||
raise ValueError(f"群 {group_id}: 无法获取平台 ID,拒绝存储消息")
|
||||
|
||||
# 5. 提取消息内容
|
||||
message_parts = self._extract_message_parts(event)
|
||||
if not message_parts:
|
||||
# 尝试记录一条警告但不中断流程(或者视为错误)
|
||||
# 原逻辑是抛出 RuntimeError
|
||||
raise RuntimeError(
|
||||
f"群 {group_id}: 消息内容为空 (sender={sender_name}),拒绝存储"
|
||||
)
|
||||
|
||||
# 6. 提取事件消息 ID 和事件时间
|
||||
msg_obj = getattr(event, "message_obj", None)
|
||||
event_message_id = str(getattr(msg_obj, "message_id", "") or "")
|
||||
|
||||
platform_name = str(event.get_platform_name() or "").strip().lower()
|
||||
reserved_event_id = False
|
||||
if platform_name in {"qq_official", "qq_official_webhook"} and event_message_id:
|
||||
reserved_event_id = self._reserve_event_id(event_message_id)
|
||||
if not reserved_event_id:
|
||||
logger.debug("[QQOfficial] 跳过重复消息事件: %s", event_message_id)
|
||||
return False
|
||||
history_content = {
|
||||
"type": "user",
|
||||
"message": message_parts,
|
||||
}
|
||||
if platform_name in {"qq_official", "qq_official_webhook"}:
|
||||
event_timestamp = self._extract_event_timestamp(msg_obj)
|
||||
history_content["_qq_official"] = {
|
||||
"message_id": event_message_id,
|
||||
"timestamp": event_timestamp,
|
||||
}
|
||||
|
||||
# 7. 存储到数据库
|
||||
try:
|
||||
await self._insert_message_history(
|
||||
platform_id=platform_id,
|
||||
group_id=group_id,
|
||||
content=history_content,
|
||||
sender_id=sender_id,
|
||||
sender_name=sender_name,
|
||||
)
|
||||
except BaseException:
|
||||
if reserved_event_id:
|
||||
self._release_event_id(event_message_id)
|
||||
raise
|
||||
else:
|
||||
if reserved_event_id:
|
||||
self._commit_event_id(event_message_id)
|
||||
|
||||
# Register the group so the scheduler can discover platforms that
|
||||
# do not provide a group-list API (Telegram, QQ Official, etc.).
|
||||
try:
|
||||
await self.group_registry.upsert(
|
||||
platform_id=platform_id,
|
||||
group_id=group_id,
|
||||
sender_id=sender_id,
|
||||
sender_name=sender_name,
|
||||
event_message_id=event_message_id,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
"[GroupRegistry] Upsert failed: "
|
||||
f"platform_id={platform_id} group_id={group_id} error={e}"
|
||||
)
|
||||
|
||||
logger.debug(
|
||||
f"[{platform_id}] 已缓存群 {group_id} 的消息 (发送者: {sender_name})"
|
||||
)
|
||||
return True
|
||||
|
||||
async def _insert_message_history(
|
||||
self,
|
||||
platform_id: str,
|
||||
group_id: str,
|
||||
content: dict,
|
||||
sender_id: str,
|
||||
sender_name: str,
|
||||
) -> None:
|
||||
"""兼容不同 AstrBot 版本的消息历史写入接口。
|
||||
|
||||
Args:
|
||||
platform_id: AstrBot 平台实例 ID。
|
||||
group_id: 当前群组 ID。
|
||||
content: 待持久化的标准化消息内容。
|
||||
sender_id: 发送者 ID。
|
||||
sender_name: 发送者展示名称。
|
||||
|
||||
Raises:
|
||||
Exception: 消息历史管理器写入失败时原样抛出。
|
||||
"""
|
||||
insert = self.context.message_history_manager.insert
|
||||
insert_kwargs = {
|
||||
"platform_id": platform_id,
|
||||
"user_id": group_id,
|
||||
"content": content,
|
||||
"sender_id": sender_id,
|
||||
"sender_name": sender_name,
|
||||
}
|
||||
if self._supports_history_max_messages is False:
|
||||
await insert(**insert_kwargs)
|
||||
return
|
||||
|
||||
try:
|
||||
await insert(**insert_kwargs, max_messages=_LOCAL_HISTORY_MAX_MESSAGES)
|
||||
except TypeError as exc:
|
||||
if "unexpected keyword argument 'max_messages'" not in str(exc):
|
||||
raise
|
||||
|
||||
self._supports_history_max_messages = False
|
||||
logger.warning(
|
||||
"[消息历史] 当前 AstrBot 核心不支持 max_messages 参数,"
|
||||
"将使用兼容模式写入消息历史;建议升级核心以启用本地历史上限。"
|
||||
)
|
||||
await insert(**insert_kwargs)
|
||||
else:
|
||||
self._supports_history_max_messages = True
|
||||
|
||||
def _get_group_id_from_event(self, event: AstrMessageEvent) -> str | None:
|
||||
"""从消息事件中安全获取群组 ID"""
|
||||
try:
|
||||
group_id = event.get_group_id()
|
||||
return group_id if group_id else None
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
def _resolve_sender_name(self, event: AstrMessageEvent, sender_id: str) -> str:
|
||||
"""解析发送者展示名"""
|
||||
platform_name = str(event.get_platform_name() or "").lower()
|
||||
candidates: list[str | None] = []
|
||||
|
||||
msg_obj = getattr(event, "message_obj", None)
|
||||
sender_obj = getattr(msg_obj, "sender", None)
|
||||
raw_message = getattr(msg_obj, "raw_message", None)
|
||||
raw_msg_obj = getattr(raw_message, "message", raw_message)
|
||||
from_user = getattr(raw_msg_obj, "from_user", None)
|
||||
|
||||
if platform_name == "telegram":
|
||||
if from_user is not None:
|
||||
candidates.extend(
|
||||
[
|
||||
getattr(from_user, "full_name", None),
|
||||
getattr(from_user, "first_name", None),
|
||||
]
|
||||
)
|
||||
candidates.append(event.get_sender_name())
|
||||
if sender_obj is not None:
|
||||
candidates.append(getattr(sender_obj, "nickname", None))
|
||||
if from_user is not None:
|
||||
candidates.append(getattr(from_user, "username", None))
|
||||
else:
|
||||
candidates.append(event.get_sender_name())
|
||||
if sender_obj is not None:
|
||||
candidates.append(getattr(sender_obj, "nickname", None))
|
||||
|
||||
if from_user is not None:
|
||||
candidates.extend(
|
||||
[
|
||||
getattr(from_user, "full_name", None),
|
||||
getattr(from_user, "first_name", None),
|
||||
getattr(from_user, "username", None),
|
||||
]
|
||||
)
|
||||
|
||||
for candidate in candidates:
|
||||
name = str(candidate or "").strip()
|
||||
if not self._is_placeholder_sender_name(name, sender_id):
|
||||
return name
|
||||
|
||||
return sender_id
|
||||
|
||||
def _extract_message_parts(self, event: AstrMessageEvent) -> list[dict]:
|
||||
"""从事件中提取消息内容"""
|
||||
message_parts = []
|
||||
message = event.message_obj
|
||||
platform_name = str(event.get_platform_name() or "").strip().lower()
|
||||
qq_mention_replacements = (
|
||||
self._extract_qq_official_mention_replacements(event)
|
||||
if platform_name in _QQ_OFFICIAL_PLATFORM_NAMES
|
||||
else None
|
||||
)
|
||||
|
||||
# 收集 @ 标记
|
||||
pending_mentions: Counter[str] = Counter()
|
||||
if message and hasattr(message, "message"):
|
||||
for seg in message.message:
|
||||
if not hasattr(seg, "type"):
|
||||
continue
|
||||
if seg.type not in ("At", "at"):
|
||||
continue
|
||||
|
||||
target = getattr(seg, "target", None) or getattr(seg, "qq", None)
|
||||
if target is None and hasattr(seg, "data"):
|
||||
target = seg.data.get("qq") or seg.data.get("target")
|
||||
|
||||
target_str = str(target or "").strip()
|
||||
if target_str:
|
||||
pending_mentions[target_str] += 1
|
||||
|
||||
display_name = str(getattr(seg, "name", "") or "").strip()
|
||||
if display_name and display_name != target_str:
|
||||
pending_mentions[display_name] += 1
|
||||
|
||||
if message and hasattr(message, "message"):
|
||||
for seg in message.message:
|
||||
if not hasattr(seg, "type"):
|
||||
continue
|
||||
|
||||
seg_type = seg.type
|
||||
if seg_type in ("Plain", "text"):
|
||||
text = getattr(seg, "text", None)
|
||||
if text is None and hasattr(seg, "data"):
|
||||
text = seg.data.get("text")
|
||||
if text:
|
||||
text = self._strip_known_mentions(text, pending_mentions)
|
||||
if qq_mention_replacements is not None:
|
||||
text = self._sanitize_qq_official_mentions(
|
||||
text, qq_mention_replacements
|
||||
)
|
||||
message_parts.append({"type": "plain", "text": text})
|
||||
|
||||
elif seg_type in ("Image", "image"):
|
||||
url = getattr(seg, "url", None) or (
|
||||
seg.data.get("url") if hasattr(seg, "data") else None
|
||||
)
|
||||
if url:
|
||||
message_parts.append({"type": "image", "url": url})
|
||||
|
||||
elif seg_type in ("At", "at"):
|
||||
target = getattr(seg, "target", None) or getattr(seg, "qq", None)
|
||||
if target is None and hasattr(seg, "data"):
|
||||
target = seg.data.get("qq") or seg.data.get("target")
|
||||
if target:
|
||||
message_parts.append(
|
||||
{
|
||||
"type": "at",
|
||||
"target_id": str(target),
|
||||
"name": str(getattr(seg, "name", "") or ""),
|
||||
}
|
||||
)
|
||||
|
||||
elif seg_type in ("File", "file"):
|
||||
url = getattr(seg, "url", None) or getattr(seg, "file_", None)
|
||||
message_parts.append(
|
||||
{
|
||||
"type": "file",
|
||||
"url": str(url or ""),
|
||||
"name": str(getattr(seg, "name", "") or ""),
|
||||
}
|
||||
)
|
||||
|
||||
elif seg_type in ("Record", "record", "voice"):
|
||||
url = getattr(seg, "url", None) or getattr(seg, "file", None)
|
||||
message_parts.append({"type": "voice", "url": str(url or "")})
|
||||
|
||||
elif seg_type in ("Video", "video"):
|
||||
url = getattr(seg, "url", None) or getattr(seg, "file", None)
|
||||
message_parts.append({"type": "video", "url": str(url or "")})
|
||||
|
||||
if not message_parts and event.message_str:
|
||||
fallback_text = str(event.message_str)
|
||||
if qq_mention_replacements is not None:
|
||||
fallback_text = self._sanitize_qq_official_mentions(
|
||||
fallback_text, qq_mention_replacements
|
||||
)
|
||||
message_parts.append({"type": "plain", "text": fallback_text})
|
||||
|
||||
# 清理空文本段
|
||||
message_parts = [
|
||||
part
|
||||
for part in message_parts
|
||||
if not (
|
||||
part.get("type") == "plain" and not str(part.get("text", "")).strip()
|
||||
)
|
||||
]
|
||||
|
||||
return message_parts
|
||||
|
||||
@classmethod
|
||||
def _extract_qq_official_mention_replacements(
|
||||
cls, event: AstrMessageEvent
|
||||
) -> dict[str, str]:
|
||||
message_obj = getattr(event, "message_obj", None)
|
||||
raw_message = getattr(message_obj, "raw_message", None)
|
||||
raw_candidates = [raw_message]
|
||||
nested_message = cls._read_field(raw_message, "message")
|
||||
if nested_message is not None and nested_message is not raw_message:
|
||||
raw_candidates.insert(0, nested_message)
|
||||
|
||||
mentions = None
|
||||
for candidate in raw_candidates:
|
||||
mentions = cls._read_field(candidate, "mentions")
|
||||
if mentions is not None:
|
||||
break
|
||||
|
||||
replacements: dict[str, str] = {}
|
||||
if not isinstance(mentions, (list, tuple)):
|
||||
return replacements
|
||||
|
||||
for mention in mentions:
|
||||
mention_id = str(
|
||||
cls._read_field(
|
||||
mention,
|
||||
"id",
|
||||
"member_openid",
|
||||
"memberopenid",
|
||||
"user_openid",
|
||||
"useropenid",
|
||||
)
|
||||
or ""
|
||||
).strip()
|
||||
if not mention_id:
|
||||
continue
|
||||
|
||||
if cls._read_field(mention, "is_you") is True:
|
||||
replacements[mention_id] = ""
|
||||
continue
|
||||
|
||||
display_name = str(
|
||||
cls._read_field(mention, "username", "name", "nickname") or ""
|
||||
).strip()
|
||||
display_name = display_name.lstrip("@").strip()
|
||||
if cls._is_placeholder_sender_name(display_name, mention_id):
|
||||
display_name = "群友"
|
||||
replacements[mention_id] = f"@{display_name}"
|
||||
|
||||
return replacements
|
||||
|
||||
@staticmethod
|
||||
def _sanitize_qq_official_mentions(text: str, replacements: dict[str, str]) -> str:
|
||||
def replace_mention(match: re.Match[str]) -> str:
|
||||
mention_id = match.group(1)
|
||||
if mention_id.lower() in {"all", "everyone"}:
|
||||
return "@全体成员"
|
||||
return replacements.get(mention_id, "@群友")
|
||||
|
||||
cleaned = _QQ_OFFICIAL_MENTION_PATTERN.sub(replace_mention, str(text))
|
||||
return re.sub(r"[^\S\r\n]{2,}", " ", cleaned).strip(" \t")
|
||||
|
||||
@staticmethod
|
||||
def _read_field(source: object, *names: str) -> object | None:
|
||||
if isinstance(source, dict):
|
||||
for name in names:
|
||||
if name in source:
|
||||
return source[name]
|
||||
return None
|
||||
|
||||
for name in names:
|
||||
value = getattr(source, name, None)
|
||||
if value is not None:
|
||||
return value
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _strip_known_mentions(text: str, pending_mentions: Counter[str]) -> str:
|
||||
"""从文本中移除已识别的 @ 提及"""
|
||||
cleaned = str(text)
|
||||
if not cleaned or not pending_mentions:
|
||||
return cleaned.strip()
|
||||
|
||||
for mention, remaining in list(pending_mentions.items()):
|
||||
if not mention or remaining <= 0:
|
||||
continue
|
||||
|
||||
pattern = re.compile(rf"(?<!\w)@{re.escape(mention)}(?!\w)")
|
||||
removed = 0
|
||||
while removed < remaining:
|
||||
cleaned, subn = pattern.subn("", cleaned, count=1)
|
||||
if subn == 0:
|
||||
break
|
||||
removed += 1
|
||||
|
||||
if removed > 0:
|
||||
pending_mentions[mention] -= removed
|
||||
if pending_mentions[mention] <= 0:
|
||||
pending_mentions.pop(mention, None)
|
||||
|
||||
return re.sub(r"[^\S\r\n]{2,}", " ", cleaned).strip()
|
||||
|
||||
@staticmethod
|
||||
def _is_placeholder_sender_name(name: str | None, sender_id: str) -> bool:
|
||||
"""判断 sender_name 是否为占位值"""
|
||||
if not name:
|
||||
return True
|
||||
normalized = str(name).strip()
|
||||
if not normalized:
|
||||
return True
|
||||
if normalized.lower() in {"unknown", "none", "null", "nil", "undefined"}:
|
||||
return True
|
||||
return normalized == str(sender_id).strip()
|
||||
|
||||
@staticmethod
|
||||
def _extract_event_timestamp(message_obj: object) -> int:
|
||||
"""从消息对象中提取平台事件时间戳。"""
|
||||
raw_message = getattr(message_obj, "raw_message", None)
|
||||
if isinstance(raw_message, dict):
|
||||
candidate = raw_message.get("timestamp")
|
||||
if not candidate:
|
||||
raw_data = raw_message.get("raw_data")
|
||||
if isinstance(raw_data, dict):
|
||||
candidate = raw_data.get("timestamp")
|
||||
else:
|
||||
raw_data = getattr(raw_message, "raw_data", None)
|
||||
candidate = getattr(raw_message, "timestamp", None)
|
||||
if not candidate and isinstance(raw_data, dict):
|
||||
candidate = raw_data.get("timestamp")
|
||||
if isinstance(candidate, (int, float)):
|
||||
return int(candidate)
|
||||
if candidate:
|
||||
try:
|
||||
from datetime import datetime
|
||||
|
||||
return int(
|
||||
datetime.fromisoformat(
|
||||
str(candidate).replace("Z", "+00:00")
|
||||
).timestamp()
|
||||
)
|
||||
except (TypeError, ValueError, OverflowError):
|
||||
pass
|
||||
return 0
|
||||
|
||||
def _reserve_event_id(self, event_message_id: str) -> bool:
|
||||
"""预占事件消息 ID:在历史记录持久化期间防止重复入库。"""
|
||||
if (
|
||||
event_message_id in self._inflight_event_ids
|
||||
or event_message_id in self._seen_event_ids
|
||||
):
|
||||
if event_message_id in self._seen_event_ids:
|
||||
self._seen_event_ids.move_to_end(event_message_id)
|
||||
return False
|
||||
self._inflight_event_ids.add(event_message_id)
|
||||
return True
|
||||
|
||||
def _commit_event_id(self, event_message_id: str) -> None:
|
||||
"""确认事件消息 ID:标记为已持久化,纳入后续去重。"""
|
||||
self._inflight_event_ids.discard(event_message_id)
|
||||
if event_message_id in self._seen_event_ids:
|
||||
self._seen_event_ids.move_to_end(event_message_id)
|
||||
else:
|
||||
self._seen_event_ids[event_message_id] = None
|
||||
if len(self._seen_event_ids) > self._seen_event_ids_limit:
|
||||
self._seen_event_ids.popitem(last=False)
|
||||
|
||||
def _release_event_id(self, event_message_id: str) -> None:
|
||||
"""释放事件消息 ID:持久化失败或取消时清理预占状态。"""
|
||||
self._inflight_event_ids.discard(event_message_id)
|
||||
@@ -0,0 +1 @@
|
||||
# 领域层 - 与平台无关的业务逻辑
|
||||
@@ -0,0 +1,18 @@
|
||||
"""
|
||||
领域实体
|
||||
|
||||
该模块导出所有领域实体类,包括:
|
||||
- AnalysisTask: 分析任务聚合根
|
||||
- IncrementalBatch: 增量分析独立批次实体
|
||||
- IncrementalState: 增量分析聚合视图(报告时使用)
|
||||
"""
|
||||
|
||||
from .analysis_task import AnalysisTask, TaskStatus
|
||||
from .incremental_state import IncrementalBatch, IncrementalState
|
||||
|
||||
__all__ = [
|
||||
"AnalysisTask",
|
||||
"TaskStatus",
|
||||
"IncrementalBatch",
|
||||
"IncrementalState",
|
||||
]
|
||||
@@ -0,0 +1,70 @@
|
||||
"""
|
||||
分析任务实体 - 聚合根
|
||||
"""
|
||||
|
||||
import time
|
||||
import uuid
|
||||
from dataclasses import dataclass, field
|
||||
from enum import Enum
|
||||
|
||||
|
||||
class TaskStatus(Enum):
|
||||
PENDING = "pending"
|
||||
CHECKING_PLATFORM = "checking_platform"
|
||||
FETCHING_MESSAGES = "fetching_messages"
|
||||
ANALYZING = "analyzing"
|
||||
GENERATING_REPORT = "generating_report"
|
||||
SENDING = "sending"
|
||||
COMPLETED = "completed"
|
||||
FAILED = "failed"
|
||||
UNSUPPORTED_PLATFORM = "unsupported_platform"
|
||||
|
||||
|
||||
@dataclass
|
||||
class AnalysisTask:
|
||||
"""分析任务实体 - 聚合根"""
|
||||
|
||||
id: str = field(default_factory=lambda: uuid.uuid4().hex[:8])
|
||||
group_id: str = ""
|
||||
platform_name: str = ""
|
||||
trace_id: str = ""
|
||||
status: TaskStatus = TaskStatus.PENDING
|
||||
is_manual: bool = False
|
||||
created_at: float = field(default_factory=time.time)
|
||||
started_at: float | None = None
|
||||
completed_at: float | None = None
|
||||
result_id: str | None = None
|
||||
error_message: str | None = None
|
||||
|
||||
def start(self, can_analyze: bool) -> bool:
|
||||
"""启动任务,验证平台能力"""
|
||||
if not can_analyze:
|
||||
self.status = TaskStatus.UNSUPPORTED_PLATFORM
|
||||
self.error_message = f"平台 {self.platform_name} 不支持分析"
|
||||
return False
|
||||
self.status = TaskStatus.FETCHING_MESSAGES
|
||||
self.started_at = time.time()
|
||||
return True
|
||||
|
||||
def advance_to(self, status: TaskStatus):
|
||||
"""推进到下一个状态"""
|
||||
self.status = status
|
||||
|
||||
def complete(self, result_id: str):
|
||||
"""标记任务为已完成"""
|
||||
self.status = TaskStatus.COMPLETED
|
||||
self.result_id = result_id
|
||||
self.completed_at = time.time()
|
||||
|
||||
def fail(self, error: str):
|
||||
"""标记任务为失败"""
|
||||
self.status = TaskStatus.FAILED
|
||||
self.error_message = error
|
||||
self.completed_at = time.time()
|
||||
|
||||
@property
|
||||
def duration(self) -> float | None:
|
||||
"""获取任务持续时间(秒)"""
|
||||
if self.started_at and self.completed_at:
|
||||
return self.completed_at - self.started_at
|
||||
return None
|
||||
@@ -0,0 +1,394 @@
|
||||
"""
|
||||
增量分析实体 — 滑动窗口批次存储架构
|
||||
|
||||
核心概念:
|
||||
- IncrementalBatch: 单次增量分析产生的独立批次数据,按批次独立存储
|
||||
- IncrementalState: 报告生成时由多个批次合并而成的聚合视图(不再持久化)
|
||||
|
||||
滑动窗口设计:
|
||||
- 每次增量分析产生一个 IncrementalBatch,独立存储到 KV
|
||||
- 最终报告时按 analysis_days × 24h 的时间窗口查询批次并合并
|
||||
- 支持同一天多次发送报告,每次都基于当前时间窗口内的所有批次
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
import uuid
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
|
||||
|
||||
@dataclass
|
||||
class IncrementalBatch:
|
||||
"""
|
||||
单次增量分析批次数据
|
||||
|
||||
每次增量分析执行完毕后产生一个 IncrementalBatch,
|
||||
包含该批次的所有统计数据和 LLM 分析结果,独立存储到 KV。
|
||||
|
||||
Attributes:
|
||||
group_id: 群组 ID
|
||||
batch_id: 批次唯一标识(UUID)
|
||||
timestamp: 批次创建时间戳(epoch)
|
||||
messages_count: 本批次分析的消息数量
|
||||
characters_count: 本批次的总字符数
|
||||
hourly_msg_counts: 按小时的消息计数 {hour_str: count}
|
||||
hourly_char_counts: 按小时的字符计数 {hour_str: count}
|
||||
user_stats: 用户统计 {user_id: {name, message_count, char_count, ...}}
|
||||
emoji_stats: 表情统计 {emoji_type: count}
|
||||
topics: 本批次提取的话题列表
|
||||
golden_quotes: 本批次提取的金句列表
|
||||
token_usage: 本批次 token 消耗 {prompt_tokens, completion_tokens, total_tokens}
|
||||
chat_quality_review: 本批次提取的聊天质量锐评
|
||||
last_message_timestamp: 本批次最后一条消息的时间戳
|
||||
participant_ids: 本批次参与者 ID 列表
|
||||
"""
|
||||
|
||||
group_id: str = ""
|
||||
batch_id: str = field(default_factory=lambda: str(uuid.uuid4()))
|
||||
timestamp: float = field(default_factory=time.time)
|
||||
|
||||
# 统计数据
|
||||
messages_count: int = 0
|
||||
characters_count: int = 0
|
||||
hourly_msg_counts: dict[str, int] = field(default_factory=dict)
|
||||
hourly_char_counts: dict[str, int] = field(default_factory=dict)
|
||||
|
||||
# 用户活跃数据
|
||||
user_stats: dict[str, dict] = field(default_factory=dict)
|
||||
|
||||
# 表情统计
|
||||
emoji_stats: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
# LLM 分析结果
|
||||
topics: list[dict] = field(default_factory=list)
|
||||
golden_quotes: list[dict] = field(default_factory=list)
|
||||
|
||||
# Token 消耗
|
||||
token_usage: dict = field(
|
||||
default_factory=lambda: {
|
||||
"prompt_tokens": 0,
|
||||
"completion_tokens": 0,
|
||||
"total_tokens": 0,
|
||||
}
|
||||
)
|
||||
|
||||
# 增量追踪
|
||||
chat_quality_review: dict[str, Any] | None = None
|
||||
last_message_timestamp: int = 0
|
||||
participant_ids: list[str] = field(default_factory=list)
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
"""序列化为字典,用于 KV 存储"""
|
||||
return {
|
||||
"group_id": self.group_id,
|
||||
"batch_id": self.batch_id,
|
||||
"timestamp": self.timestamp,
|
||||
"messages_count": self.messages_count,
|
||||
"characters_count": self.characters_count,
|
||||
"hourly_msg_counts": self.hourly_msg_counts,
|
||||
"hourly_char_counts": self.hourly_char_counts,
|
||||
"user_stats": self.user_stats,
|
||||
"emoji_stats": self.emoji_stats,
|
||||
"topics": self.topics,
|
||||
"golden_quotes": self.golden_quotes,
|
||||
"token_usage": self.token_usage,
|
||||
"chat_quality_review": self.chat_quality_review,
|
||||
"last_message_timestamp": self.last_message_timestamp,
|
||||
"participant_ids": self.participant_ids,
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, data: dict) -> IncrementalBatch:
|
||||
"""从字典反序列化"""
|
||||
return cls(
|
||||
group_id=data.get("group_id", ""),
|
||||
batch_id=data.get("batch_id", ""),
|
||||
timestamp=data.get("timestamp", 0.0),
|
||||
messages_count=data.get("messages_count", 0),
|
||||
characters_count=data.get("characters_count", 0),
|
||||
hourly_msg_counts=data.get("hourly_msg_counts", {}),
|
||||
hourly_char_counts=data.get("hourly_char_counts", {}),
|
||||
user_stats=data.get("user_stats", {}),
|
||||
emoji_stats=data.get("emoji_stats", {}),
|
||||
topics=data.get("topics", []),
|
||||
golden_quotes=data.get("golden_quotes", []),
|
||||
token_usage=data.get(
|
||||
"token_usage",
|
||||
{
|
||||
"prompt_tokens": 0,
|
||||
"completion_tokens": 0,
|
||||
"total_tokens": 0,
|
||||
},
|
||||
),
|
||||
chat_quality_review=data.get("chat_quality_review"),
|
||||
last_message_timestamp=data.get("last_message_timestamp", 0),
|
||||
participant_ids=data.get("participant_ids", []),
|
||||
)
|
||||
|
||||
def get_summary(self) -> dict:
|
||||
"""获取批次摘要信息"""
|
||||
return {
|
||||
"batch_id": self.batch_id[:8],
|
||||
"timestamp": datetime.fromtimestamp(self.timestamp).strftime(
|
||||
"%Y-%m-%d %H:%M:%S"
|
||||
),
|
||||
"messages_count": self.messages_count,
|
||||
"topics_count": len(self.topics),
|
||||
"quotes_count": len(self.golden_quotes),
|
||||
"participants": len(self.participant_ids),
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class IncrementalState:
|
||||
"""
|
||||
增量分析聚合视图(报告时使用)
|
||||
|
||||
由多个 IncrementalBatch 合并而成,不直接持久化。
|
||||
IncrementalMergeService.merge_batches() 负责从批次列表构建此对象。
|
||||
|
||||
Attributes:
|
||||
group_id: 群组 ID
|
||||
window_start: 滑动窗口起始时间戳
|
||||
window_end: 滑动窗口结束时间戳
|
||||
topics: 合并去重后的话题列表
|
||||
golden_quotes: 合并去重后的金句列表
|
||||
hourly_message_counts: 合并后的每小时消息计数 {hour_str: count}
|
||||
hourly_character_counts: 合并后的每小时字符计数 {hour_str: count}
|
||||
user_activities: 合并后的用户活跃数据
|
||||
emoji_counts: 合并后的表情统计
|
||||
total_message_count: 窗口内总消息数
|
||||
total_character_count: 窗口内总字符数
|
||||
total_analysis_count: 窗口内批次数量
|
||||
total_token_usage: 累计 token 消耗
|
||||
last_analyzed_message_timestamp: 最后分析消息时间戳
|
||||
all_participant_ids: 所有参与者 ID 集合
|
||||
"""
|
||||
|
||||
# 标识信息
|
||||
group_id: str = ""
|
||||
window_start: float = 0.0
|
||||
window_end: float = 0.0
|
||||
|
||||
# 合并后的 LLM 分析结果
|
||||
topics: list[dict] = field(default_factory=list)
|
||||
golden_quotes: list[dict] = field(default_factory=list)
|
||||
chat_quality_review: dict[str, Any] | None = None
|
||||
all_quality_reviews: list[dict] = field(
|
||||
default_factory=list
|
||||
) # 存储所有批次的质量锐评,用于最终报告时的汇总分析
|
||||
|
||||
# 合并后的统计数据(按小时)
|
||||
hourly_message_counts: dict[str, int] = field(default_factory=dict)
|
||||
hourly_character_counts: dict[str, int] = field(default_factory=dict)
|
||||
|
||||
# 用户活跃数据
|
||||
user_activities: dict[str, dict] = field(default_factory=dict)
|
||||
|
||||
# 表情统计
|
||||
emoji_counts: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
# 汇总统计
|
||||
total_message_count: int = 0
|
||||
total_character_count: int = 0
|
||||
total_analysis_count: int = 0
|
||||
total_token_usage: dict = field(
|
||||
default_factory=lambda: {
|
||||
"prompt_tokens": 0,
|
||||
"completion_tokens": 0,
|
||||
"total_tokens": 0,
|
||||
}
|
||||
)
|
||||
|
||||
# 增量跟踪
|
||||
last_analyzed_message_timestamp: int = 0
|
||||
all_participant_ids: set[str] = field(default_factory=set)
|
||||
|
||||
# 元数据
|
||||
created_at: float = field(default_factory=time.time)
|
||||
updated_at: float = field(default_factory=time.time)
|
||||
|
||||
def get_peak_hours(self, top_n: int = 3) -> list[int]:
|
||||
"""
|
||||
获取消息最活跃的时段。
|
||||
|
||||
Args:
|
||||
top_n: 返回前 N 个最活跃的小时
|
||||
|
||||
Returns:
|
||||
list[int]: 活跃小时列表,按消息量降序
|
||||
"""
|
||||
if not self.hourly_message_counts:
|
||||
return []
|
||||
sorted_hours = sorted(
|
||||
self.hourly_message_counts.items(),
|
||||
key=lambda x: x[1],
|
||||
reverse=True,
|
||||
)
|
||||
return [int(h) for h, _ in sorted_hours[:top_n]]
|
||||
|
||||
def get_most_active_period(self) -> str:
|
||||
"""
|
||||
获取最活跃时段的描述字符串。
|
||||
|
||||
Returns:
|
||||
str: 如 "20:00-21:00"
|
||||
"""
|
||||
peak = self.get_peak_hours(1)
|
||||
if not peak:
|
||||
return "未知"
|
||||
hour = peak[0]
|
||||
return f"{hour:02d}:00-{hour + 1:02d}:00"
|
||||
|
||||
def get_user_activity_ranking(self, top_n: int = 10) -> list[dict]:
|
||||
"""
|
||||
获取用户活跃度排名。
|
||||
|
||||
Args:
|
||||
top_n: 返回前 N 名
|
||||
|
||||
Returns:
|
||||
list[dict]: 按消息数降序排列的用户列表
|
||||
"""
|
||||
users = []
|
||||
for user_id, data in self.user_activities.items():
|
||||
users.append(
|
||||
{
|
||||
"user_id": user_id,
|
||||
"name": data.get("nickname", data.get("name", user_id)),
|
||||
"message_count": data.get("message_count", 0),
|
||||
"char_count": data.get("char_count", 0),
|
||||
}
|
||||
)
|
||||
users.sort(key=lambda x: x["message_count"], reverse=True)
|
||||
return users[:top_n]
|
||||
|
||||
def get_window_date_str(self) -> str:
|
||||
"""
|
||||
获取窗口的日期范围字符串,用于报告显示。
|
||||
|
||||
Returns:
|
||||
str: 如 "2024-01-15" 或 "2024-01-14 ~ 2024-01-15"
|
||||
"""
|
||||
if self.window_start <= 0 or self.window_end <= 0:
|
||||
return datetime.now().strftime("%Y-%m-%d")
|
||||
|
||||
start_date = datetime.fromtimestamp(self.window_start).strftime("%Y-%m-%d")
|
||||
end_date = datetime.fromtimestamp(self.window_end).strftime("%Y-%m-%d")
|
||||
|
||||
if start_date == end_date:
|
||||
return end_date
|
||||
return f"{start_date} ~ {end_date}"
|
||||
|
||||
def get_summary(self) -> dict:
|
||||
"""
|
||||
获取当前增量状态的摘要信息,用于状态查询命令。
|
||||
|
||||
Returns:
|
||||
dict: 包含关键统计信息的摘要
|
||||
"""
|
||||
return {
|
||||
"group_id": self.group_id,
|
||||
"window": self.get_window_date_str(),
|
||||
"total_messages": self.total_message_count,
|
||||
"total_characters": self.total_character_count,
|
||||
"total_analyses": self.total_analysis_count,
|
||||
"topics_count": len(self.topics),
|
||||
"quotes_count": len(self.golden_quotes),
|
||||
"participants": len(self.all_participant_ids),
|
||||
"total_tokens": self.total_token_usage.get("total_tokens", 0),
|
||||
"last_analysis_time": (
|
||||
datetime.fromtimestamp(self.updated_at).strftime("%H:%M:%S")
|
||||
if self.updated_at
|
||||
else "无"
|
||||
),
|
||||
"peak_hours": self.get_peak_hours(3),
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def is_duplicate_topic(
|
||||
new_topic: dict, existing_topics: list[dict], threshold: float = 0.6
|
||||
) -> bool:
|
||||
"""
|
||||
检测话题是否与已有话题重复。
|
||||
|
||||
使用简单的字符重叠相似度判断。
|
||||
当新话题的名称与已有话题名称相似度超过阈值时,认为是重复话题。
|
||||
|
||||
Args:
|
||||
new_topic: 待检测的新话题
|
||||
existing_topics: 已有话题列表
|
||||
threshold: 相似度阈值(0-1),默认 0.6
|
||||
|
||||
Returns:
|
||||
bool: 是否重复
|
||||
"""
|
||||
new_name = new_topic.get("topic", "")
|
||||
if not new_name:
|
||||
return False
|
||||
|
||||
for existing in existing_topics:
|
||||
existing_name = existing.get("topic", "")
|
||||
if not existing_name:
|
||||
continue
|
||||
similarity = IncrementalState.char_overlap_similarity(
|
||||
new_name, existing_name
|
||||
)
|
||||
if similarity >= threshold:
|
||||
return True
|
||||
return False
|
||||
|
||||
@staticmethod
|
||||
def is_duplicate_quote(
|
||||
new_quote: dict, existing_quotes: list[dict], threshold: float = 0.7
|
||||
) -> bool:
|
||||
"""
|
||||
检测金句是否与已有金句重复。
|
||||
|
||||
Args:
|
||||
new_quote: 待检测的新金句
|
||||
existing_quotes: 已有金句列表
|
||||
threshold: 相似度阈值(0-1),默认 0.7
|
||||
|
||||
Returns:
|
||||
bool: 是否重复
|
||||
"""
|
||||
new_content = new_quote.get("content", "")
|
||||
if not new_content:
|
||||
return False
|
||||
|
||||
for existing in existing_quotes:
|
||||
existing_content = existing.get("content", "")
|
||||
if not existing_content:
|
||||
continue
|
||||
similarity = IncrementalState.char_overlap_similarity(
|
||||
new_content, existing_content
|
||||
)
|
||||
if similarity >= threshold:
|
||||
return True
|
||||
return False
|
||||
|
||||
@staticmethod
|
||||
def char_overlap_similarity(s1: str, s2: str) -> float:
|
||||
"""
|
||||
计算两个字符串的字符重叠相似度(Jaccard 相似系数)。
|
||||
|
||||
Args:
|
||||
s1: 第一个字符串
|
||||
s2: 第二个字符串
|
||||
|
||||
Returns:
|
||||
float: 相似度值(0-1)
|
||||
"""
|
||||
if not s1 or not s2:
|
||||
return 0.0
|
||||
set1 = set(s1)
|
||||
set2 = set(s2)
|
||||
intersection = set1 & set2
|
||||
union = set1 | set2
|
||||
if not union:
|
||||
return 0.0
|
||||
return len(intersection) / len(union)
|
||||
@@ -0,0 +1,254 @@
|
||||
"""
|
||||
领域异常 - 领域层自定义异常
|
||||
|
||||
该模块包含插件中使用的所有领域特定异常。
|
||||
这些异常是平台无关的,表示业务逻辑错误。
|
||||
"""
|
||||
|
||||
|
||||
class DomainException(Exception):
|
||||
"""所有领域错误的基础异常。"""
|
||||
|
||||
def __init__(self, message: str, code: str = "DOMAIN_ERROR"):
|
||||
self.message = message
|
||||
self.code = code
|
||||
super().__init__(self.message)
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# 分析异常
|
||||
# ============================================================================
|
||||
|
||||
|
||||
class AnalysisException(DomainException):
|
||||
"""分析相关错误的基础异常。"""
|
||||
|
||||
def __init__(self, message: str, code: str = "ANALYSIS_ERROR"):
|
||||
super().__init__(message, code)
|
||||
|
||||
|
||||
class InsufficientDataException(AnalysisException):
|
||||
"""当数据不足以进行分析时抛出。"""
|
||||
|
||||
def __init__(self, message: str = "数据不足,无法进行分析"):
|
||||
super().__init__(message, "INSUFFICIENT_DATA")
|
||||
|
||||
|
||||
class AnalysisTimeoutException(AnalysisException):
|
||||
"""当分析超时时抛出。"""
|
||||
|
||||
def __init__(self, message: str = "分析超时"):
|
||||
super().__init__(message, "ANALYSIS_TIMEOUT")
|
||||
|
||||
|
||||
class LLMException(AnalysisException):
|
||||
"""当 LLM API 调用失败时抛出。"""
|
||||
|
||||
def __init__(self, message: str = "LLM API 调用失败", provider: str = ""):
|
||||
self.provider = provider
|
||||
super().__init__(
|
||||
f"{message} (提供商: {provider})" if provider else message, "LLM_ERROR"
|
||||
)
|
||||
|
||||
|
||||
class LLMRateLimitException(LLMException):
|
||||
"""当 LLM API 速率限制超出时抛出。"""
|
||||
|
||||
def __init__(self, message: str = "LLM 速率限制超出", provider: str = ""):
|
||||
super().__init__(message, provider)
|
||||
self.code = "LLM_RATE_LIMIT"
|
||||
|
||||
|
||||
class LLMQuotaExceededException(LLMException):
|
||||
"""当 LLM API 配额超出时抛出。"""
|
||||
|
||||
def __init__(self, message: str = "LLM 配额超出", provider: str = ""):
|
||||
super().__init__(message, provider)
|
||||
self.code = "LLM_QUOTA_EXCEEDED"
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# 平台异常
|
||||
# ============================================================================
|
||||
|
||||
|
||||
class PlatformException(DomainException):
|
||||
"""平台相关错误的基础异常。"""
|
||||
|
||||
def __init__(self, message: str, platform: str = "", code: str = "PLATFORM_ERROR"):
|
||||
self.platform = platform
|
||||
super().__init__(f"[{platform}] {message}" if platform else message, code)
|
||||
|
||||
|
||||
class PlatformNotSupportedException(PlatformException):
|
||||
"""当平台不被支持时抛出。"""
|
||||
|
||||
def __init__(self, platform: str):
|
||||
super().__init__(
|
||||
f"平台 '{platform}' 不被支持", platform, "PLATFORM_NOT_SUPPORTED"
|
||||
)
|
||||
|
||||
|
||||
class PlatformConnectionException(PlatformException):
|
||||
"""当连接平台失败时抛出。"""
|
||||
|
||||
def __init__(self, message: str = "连接平台失败", platform: str = ""):
|
||||
super().__init__(message, platform, "PLATFORM_CONNECTION_ERROR")
|
||||
|
||||
|
||||
class PlatformAPIException(PlatformException):
|
||||
"""当平台 API 调用失败时抛出。"""
|
||||
|
||||
def __init__(self, message: str = "平台 API 调用失败", platform: str = ""):
|
||||
super().__init__(message, platform, "PLATFORM_API_ERROR")
|
||||
|
||||
|
||||
class MessageFetchException(PlatformException):
|
||||
"""当获取消息失败时抛出。"""
|
||||
|
||||
def __init__(
|
||||
self, message: str = "获取消息失败", platform: str = "", group_id: str = ""
|
||||
):
|
||||
self.group_id = group_id
|
||||
super().__init__(
|
||||
f"{message} (群组: {group_id})" if group_id else message,
|
||||
platform,
|
||||
"MESSAGE_FETCH_ERROR",
|
||||
)
|
||||
|
||||
|
||||
class MessageSendException(PlatformException):
|
||||
"""当发送消息失败时抛出。"""
|
||||
|
||||
def __init__(
|
||||
self, message: str = "发送消息失败", platform: str = "", group_id: str = ""
|
||||
):
|
||||
self.group_id = group_id
|
||||
super().__init__(
|
||||
f"{message} (群组: {group_id})" if group_id else message,
|
||||
platform,
|
||||
"MESSAGE_SEND_ERROR",
|
||||
)
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# 配置异常
|
||||
# ============================================================================
|
||||
|
||||
|
||||
class ConfigurationException(DomainException):
|
||||
"""配置相关错误的基础异常。"""
|
||||
|
||||
def __init__(self, message: str, code: str = "CONFIG_ERROR"):
|
||||
super().__init__(message, code)
|
||||
|
||||
|
||||
class InvalidConfigurationException(ConfigurationException):
|
||||
"""当配置无效时抛出。"""
|
||||
|
||||
def __init__(self, message: str = "无效的配置", key: str = ""):
|
||||
self.key = key
|
||||
super().__init__(f"{message}: {key}" if key else message, "INVALID_CONFIG")
|
||||
|
||||
|
||||
class MissingConfigurationException(ConfigurationException):
|
||||
"""当缺少必需配置时抛出。"""
|
||||
|
||||
def __init__(self, key: str):
|
||||
self.key = key
|
||||
super().__init__(f"缺少必需配置: {key}", "MISSING_CONFIG")
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# 仓储异常
|
||||
# ============================================================================
|
||||
|
||||
|
||||
class RepositoryException(DomainException):
|
||||
"""仓储相关错误的基础异常。"""
|
||||
|
||||
def __init__(self, message: str, code: str = "REPOSITORY_ERROR"):
|
||||
super().__init__(message, code)
|
||||
|
||||
|
||||
class DataNotFoundException(RepositoryException):
|
||||
"""当请求的数据未找到时抛出。"""
|
||||
|
||||
def __init__(
|
||||
self, message: str = "数据未找到", entity_type: str = "", entity_id: str = ""
|
||||
):
|
||||
self.entity_type = entity_type
|
||||
self.entity_id = entity_id
|
||||
super().__init__(
|
||||
f"{entity_type} 未找到: {entity_id}" if entity_type else message,
|
||||
"DATA_NOT_FOUND",
|
||||
)
|
||||
|
||||
|
||||
class DataPersistenceException(RepositoryException):
|
||||
"""当数据持久化失败时抛出。"""
|
||||
|
||||
def __init__(self, message: str = "数据持久化失败"):
|
||||
super().__init__(message, "DATA_PERSISTENCE_ERROR")
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# 调度异常
|
||||
# ============================================================================
|
||||
|
||||
|
||||
class SchedulingException(DomainException):
|
||||
"""调度相关错误的基础异常。"""
|
||||
|
||||
def __init__(self, message: str, code: str = "SCHEDULING_ERROR"):
|
||||
super().__init__(message, code)
|
||||
|
||||
|
||||
class TaskAlreadyScheduledException(SchedulingException):
|
||||
"""当尝试调度已调度的任务时抛出。"""
|
||||
|
||||
def __init__(self, task_id: str):
|
||||
self.task_id = task_id
|
||||
super().__init__(f"任务已调度: {task_id}", "TASK_ALREADY_SCHEDULED")
|
||||
|
||||
|
||||
class TaskNotFoundException(SchedulingException):
|
||||
"""当找不到已调度的任务时抛出。"""
|
||||
|
||||
def __init__(self, task_id: str):
|
||||
self.task_id = task_id
|
||||
super().__init__(f"未找到已调度的任务: {task_id}", "TASK_NOT_FOUND")
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# 验证异常
|
||||
# ============================================================================
|
||||
|
||||
|
||||
class ValidationException(DomainException):
|
||||
"""验证错误的基础异常。"""
|
||||
|
||||
def __init__(self, message: str, field: str = "", code: str = "VALIDATION_ERROR"):
|
||||
self.field = field
|
||||
super().__init__(f"{field}: {message}" if field else message, code)
|
||||
|
||||
|
||||
class InvalidGroupIdException(ValidationException):
|
||||
"""当群组 ID 无效时抛出。"""
|
||||
|
||||
def __init__(self, group_id: str):
|
||||
super().__init__(f"无效的群组 ID: {group_id}", "group_id", "INVALID_GROUP_ID")
|
||||
|
||||
|
||||
class InvalidUserIdException(ValidationException):
|
||||
"""当用户 ID 无效时抛出。"""
|
||||
|
||||
def __init__(self, user_id: str):
|
||||
super().__init__(f"无效的用户 ID: {user_id}", "user_id", "INVALID_USER_ID")
|
||||
|
||||
|
||||
class InvalidMessageException(ValidationException):
|
||||
"""当消息格式无效时抛出。"""
|
||||
|
||||
def __init__(self, message: str = "无效的消息格式"):
|
||||
super().__init__(message, "message", "INVALID_MESSAGE")
|
||||
@@ -0,0 +1,121 @@
|
||||
"""
|
||||
数据模型定义
|
||||
包含所有分析相关的数据结构
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Optional
|
||||
|
||||
|
||||
@dataclass
|
||||
class SummaryTopic:
|
||||
"""话题总结数据结构"""
|
||||
|
||||
topic: str
|
||||
contributors: list[str]
|
||||
detail: str
|
||||
contributor_ids: list[str] = field(
|
||||
default_factory=list
|
||||
) # 贡献者ID列表 (用于显示头像)
|
||||
|
||||
|
||||
@dataclass
|
||||
class UserTitle:
|
||||
"""用户称号数据结构"""
|
||||
|
||||
name: str
|
||||
user_id: str # 原 qq 字段
|
||||
title: str
|
||||
mbti: str
|
||||
reason: str
|
||||
|
||||
|
||||
@dataclass
|
||||
class GoldenQuote:
|
||||
"""群聊金句数据结构"""
|
||||
|
||||
content: str
|
||||
sender: str
|
||||
reason: str
|
||||
user_id: str = "" # 原 qq 字段
|
||||
|
||||
|
||||
@dataclass
|
||||
class QualityDimension:
|
||||
"""聊天质量维度数据结构"""
|
||||
|
||||
name: str # 维度名称
|
||||
percentage: float # 占比
|
||||
comment: str # 犀利点评
|
||||
color: str = "#607d8b" # 颜色
|
||||
|
||||
|
||||
@dataclass
|
||||
class QualityReview:
|
||||
"""聊天质量锐评数据结构"""
|
||||
|
||||
title: str
|
||||
subtitle: str
|
||||
dimensions: list[QualityDimension]
|
||||
summary: str
|
||||
|
||||
|
||||
@dataclass
|
||||
class TokenUsage:
|
||||
"""Token使用统计"""
|
||||
|
||||
prompt_tokens: int = 0
|
||||
completion_tokens: int = 0
|
||||
total_tokens: int = 0
|
||||
|
||||
|
||||
@dataclass
|
||||
class EmojiStatistics:
|
||||
"""表情统计数据结构"""
|
||||
|
||||
face_count: int = 0 # QQ基础表情数量
|
||||
mface_count: int = 0 # 动画表情数量
|
||||
bface_count: int = 0 # 超级表情数量
|
||||
sface_count: int = 0 # 小表情数量
|
||||
other_emoji_count: int = 0 # 其他表情数量
|
||||
face_details: dict = field(default_factory=dict) # 具体表情ID统计 {face_id: count}
|
||||
|
||||
@property
|
||||
def total_emoji_count(self) -> int:
|
||||
"""总表情数量"""
|
||||
return (
|
||||
self.face_count
|
||||
+ self.mface_count
|
||||
+ self.bface_count
|
||||
+ self.sface_count
|
||||
+ self.other_emoji_count
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ActivityVisualization:
|
||||
"""活跃度可视化数据结构"""
|
||||
|
||||
hourly_activity: dict = field(default_factory=dict) # {hour: count}
|
||||
daily_activity: dict = field(default_factory=dict) # {date: count}
|
||||
user_activity_ranking: list = field(default_factory=list) # 用户活跃度排行
|
||||
peak_hours: list = field(default_factory=list) # 高峰时段
|
||||
activity_heatmap_data: dict = field(default_factory=dict) # 热力图数据
|
||||
|
||||
|
||||
@dataclass
|
||||
class GroupStatistics:
|
||||
"""群聊统计数据结构"""
|
||||
|
||||
message_count: int
|
||||
total_characters: int
|
||||
participant_count: int
|
||||
most_active_period: str
|
||||
golden_quotes: list[GoldenQuote]
|
||||
emoji_count: int # 保持向后兼容
|
||||
emoji_statistics: EmojiStatistics = field(default_factory=EmojiStatistics)
|
||||
activity_visualization: ActivityVisualization = field(
|
||||
default_factory=ActivityVisualization
|
||||
)
|
||||
token_usage: TokenUsage = field(default_factory=TokenUsage)
|
||||
chat_quality_review: Optional["QualityReview"] = None
|
||||
@@ -0,0 +1,12 @@
|
||||
# 仓储接口
|
||||
from .avatar_repository import IAvatarRepository
|
||||
from .message_repository import IGroupInfoRepository, IMessageRepository, IMessageSender
|
||||
from .visualization_repository import IActivityVisualizer
|
||||
|
||||
__all__ = [
|
||||
"IMessageRepository",
|
||||
"IMessageSender",
|
||||
"IGroupInfoRepository",
|
||||
"IAvatarRepository",
|
||||
"IActivityVisualizer",
|
||||
]
|
||||
@@ -0,0 +1,97 @@
|
||||
"""
|
||||
分析服务接口 - 领域层
|
||||
定义语义分析的抽象契约
|
||||
"""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
|
||||
from ..models.data_models import (
|
||||
GoldenQuote,
|
||||
QualityReview,
|
||||
SummaryTopic,
|
||||
TokenUsage,
|
||||
UserTitle,
|
||||
)
|
||||
|
||||
|
||||
class IAnalysisProvider(ABC):
|
||||
"""
|
||||
LLM 分析提供商接口
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
async def analyze_topics(
|
||||
self,
|
||||
messages: list[dict],
|
||||
umo: str | None = None,
|
||||
session_id: str | None = None,
|
||||
) -> tuple[list[SummaryTopic], TokenUsage]:
|
||||
"""分析话题"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def analyze_user_titles(
|
||||
self,
|
||||
messages: list[dict],
|
||||
user_activity: dict,
|
||||
umo: str | None = None,
|
||||
top_users: list[dict] | None = None,
|
||||
session_id: str | None = None,
|
||||
) -> tuple[list[UserTitle], TokenUsage]:
|
||||
"""分析用户称号"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def analyze_golden_quotes(
|
||||
self,
|
||||
messages: list[dict],
|
||||
umo: str | None = None,
|
||||
session_id: str | None = None,
|
||||
) -> tuple[list[GoldenQuote], TokenUsage]:
|
||||
"""分析金句"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def analyze_all_concurrent(
|
||||
self,
|
||||
messages: list[dict],
|
||||
user_activity: dict,
|
||||
umo: str | None = None,
|
||||
top_users: list[dict] | None = None,
|
||||
topic_enabled: bool = True,
|
||||
user_title_enabled: bool = True,
|
||||
golden_quote_enabled: bool = True,
|
||||
chat_quality_enabled: bool = False,
|
||||
) -> tuple[
|
||||
list[SummaryTopic],
|
||||
list[UserTitle],
|
||||
list[GoldenQuote],
|
||||
TokenUsage,
|
||||
QualityReview | None,
|
||||
]:
|
||||
"""并发分析所有内容"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def analyze_incremental_concurrent(
|
||||
self,
|
||||
messages: list[dict],
|
||||
umo: str | None = None,
|
||||
topics_per_batch: int = 3,
|
||||
quotes_per_batch: int = 3,
|
||||
topic_enabled: bool = True,
|
||||
golden_quote_enabled: bool = True,
|
||||
chat_quality_enabled: bool = False,
|
||||
) -> tuple[list[SummaryTopic], list[GoldenQuote], TokenUsage, QualityReview | None]:
|
||||
"""增量模式并发分析"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def summarize_quality_reviews(
|
||||
self,
|
||||
batch_reviews: list[dict],
|
||||
umo: str | None = None,
|
||||
session_id: str | None = None,
|
||||
) -> tuple[QualityReview | None, TokenUsage]:
|
||||
"""汇总多个聊天质量报告(增量模式使用)"""
|
||||
pass
|
||||
@@ -0,0 +1,78 @@
|
||||
"""
|
||||
头像仓储接口 - 跨平台头像抽象
|
||||
"""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
|
||||
|
||||
class IAvatarRepository(ABC):
|
||||
"""
|
||||
头像仓储接口
|
||||
|
||||
不同平台获取头像的方式不同:
|
||||
- QQ/OneBot: URL 模板 (q1.qlogo.cn)
|
||||
- Telegram: API 调用 (getUserProfilePhotos + getFile)
|
||||
- Discord: CDN URL 模板 (cdn.discordapp.com)
|
||||
- Slack: users.info API profile.image_* 字段
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
async def get_user_avatar_url(
|
||||
self,
|
||||
user_id: str,
|
||||
size: int = 100,
|
||||
) -> str | None:
|
||||
"""
|
||||
获取用户头像 URL
|
||||
|
||||
参数:
|
||||
user_id: 用户 ID
|
||||
size: 期望的头像尺寸(将选择最接近的可用尺寸)
|
||||
|
||||
返回:
|
||||
头像 URL,如果不可用则返回 None
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def get_user_avatar_data(
|
||||
self,
|
||||
user_id: str,
|
||||
size: int = 100,
|
||||
) -> str | None:
|
||||
"""
|
||||
获取用户头像的 Base64 数据
|
||||
|
||||
用于需要嵌入图片的场景(如 HTML 模板渲染)
|
||||
|
||||
返回:
|
||||
Base64 编码的图片数据 (data:image/png;base64,...),
|
||||
如果不可用则返回 None
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def get_group_avatar_url(
|
||||
self,
|
||||
group_id: str,
|
||||
size: int = 100,
|
||||
) -> str | None:
|
||||
"""获取群组头像 URL"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def batch_get_avatar_urls(
|
||||
self,
|
||||
user_ids: list[str],
|
||||
size: int = 100,
|
||||
) -> dict[str, str | None]:
|
||||
"""
|
||||
批量获取用户头像 URL
|
||||
|
||||
用于报告生成时需要一次获取多个头像
|
||||
"""
|
||||
pass
|
||||
|
||||
def get_default_avatar_url(self) -> str:
|
||||
"""获取默认头像 URL(当用户头像不可用时)"""
|
||||
return "data:image/svg+xml;base64,PHN2ZyB4bWxucz0iaHR0cDovL3d3dy53My5vcmcvMjAwMC9zdmciIHZpZXdCb3g9IjAgMCAyNCAyNCI+PHBhdGggZD0iTTEyIDEyYzIuMjEgMCA0LTEuNzkgNC00cy0xLjc5LTQtNC00LTQgMS43OS00IDQgMS43OSA0IDQgNHptMCAyYy0yLjY3IDAtOCAxLjM0LTggNHYyaDE2di0yYzAtMi42Ni01LjMzLTQtOC00eiIvPjwvc3ZnPg=="
|
||||
@@ -0,0 +1,130 @@
|
||||
"""
|
||||
消息仓储接口 - 平台无关的抽象
|
||||
"""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
|
||||
from ..value_objects.platform_capabilities import PlatformCapabilities
|
||||
from ..value_objects.unified_group import UnifiedGroup, UnifiedMember
|
||||
from ..value_objects.unified_message import UnifiedMessage
|
||||
|
||||
|
||||
class IMessageRepository(ABC):
|
||||
"""
|
||||
消息仓储接口
|
||||
|
||||
每个平台适配器必须实现此接口。
|
||||
所有方法返回统一格式,隐藏平台差异。
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
async def fetch_messages(
|
||||
self,
|
||||
group_id: str,
|
||||
days: int = 1,
|
||||
max_count: int = 1000,
|
||||
before_id: str | None = None,
|
||||
since_ts: int | None = None,
|
||||
) -> list[UnifiedMessage]:
|
||||
"""
|
||||
获取群组消息历史
|
||||
|
||||
参数:
|
||||
group_id: 群组 ID
|
||||
days: 获取最近 N 天的消息
|
||||
max_count: 最大消息数量
|
||||
before_id: 获取此 ID 之前的消息(用于分页)
|
||||
since_ts: 从指定时间戳开始拉取消息(Unix timestamp),优先级高于 days。
|
||||
|
||||
返回:
|
||||
统一消息列表,按时间升序排列
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def get_capabilities(self) -> PlatformCapabilities:
|
||||
"""获取平台能力"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def get_platform_name(self) -> str:
|
||||
"""获取平台名称"""
|
||||
pass
|
||||
|
||||
|
||||
class IMessageSender(ABC):
|
||||
"""消息发送接口"""
|
||||
|
||||
@abstractmethod
|
||||
async def send_text(
|
||||
self,
|
||||
group_id: str,
|
||||
text: str,
|
||||
reply_to: str | None = None,
|
||||
) -> bool:
|
||||
"""发送文本消息"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def send_image(
|
||||
self,
|
||||
group_id: str,
|
||||
image_path: str,
|
||||
caption: str = "",
|
||||
) -> bool:
|
||||
"""发送图片消息"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def send_forward_msg(
|
||||
self,
|
||||
group_id: str,
|
||||
nodes: list[dict],
|
||||
) -> bool:
|
||||
"""
|
||||
发送合并转发消息。
|
||||
|
||||
Args:
|
||||
group_id: 目标群组 ID
|
||||
nodes: 转发节点列表。每个节点通常包含 name, uin (或 user_id), content。
|
||||
目前主要用于 OneBot 兼容性。
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def send_file(
|
||||
self,
|
||||
group_id: str,
|
||||
file_path: str,
|
||||
filename: str | None = None,
|
||||
) -> bool:
|
||||
"""发送文件"""
|
||||
pass
|
||||
|
||||
|
||||
class IGroupInfoRepository(ABC):
|
||||
"""群组信息仓储接口"""
|
||||
|
||||
@abstractmethod
|
||||
async def get_group_info(self, group_id: str) -> UnifiedGroup | None:
|
||||
"""获取群组信息"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def get_group_list(self) -> list[str]:
|
||||
"""获取机器人所在的所有群组 ID"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def get_member_list(self, group_id: str) -> list[UnifiedMember]:
|
||||
"""获取群组成员列表"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def get_member_info(
|
||||
self,
|
||||
group_id: str,
|
||||
user_id: str,
|
||||
) -> UnifiedMember | None:
|
||||
"""获取指定成员信息"""
|
||||
pass
|
||||
@@ -0,0 +1,56 @@
|
||||
"""
|
||||
报告生成接口 - 领域层
|
||||
定义分析报告生成的抽象契约
|
||||
"""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Any
|
||||
|
||||
|
||||
class IReportGenerator(ABC):
|
||||
"""
|
||||
报告生成器接口
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
async def generate_image_report(
|
||||
self,
|
||||
analysis_result: dict,
|
||||
group_id: str,
|
||||
html_render_func: Any,
|
||||
avatar_url_getter: Any = None,
|
||||
nickname_getter: Any = None,
|
||||
avatar_cache_namespace: str | None = None,
|
||||
hide_user_names: bool = False,
|
||||
allow_alphanumeric_user_ids: bool = False,
|
||||
template_theme: str | None = None,
|
||||
) -> tuple[str | None, str | None]:
|
||||
"""生成图片报告"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def generate_html_report(
|
||||
self,
|
||||
analysis_result: dict,
|
||||
group_id: str,
|
||||
avatar_url_getter: Any = None,
|
||||
nickname_getter: Any = None,
|
||||
avatar_cache_namespace: str | None = None,
|
||||
hide_user_names: bool = False,
|
||||
allow_alphanumeric_user_ids: bool = False,
|
||||
template_theme: str | None = None,
|
||||
custom_filename: str | None = None,
|
||||
trace_id: str | None = None,
|
||||
) -> tuple[str | None, str | None]:
|
||||
"""生成 HTML 报告"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def generate_text_report(self, analysis_result: dict) -> str:
|
||||
"""生成文本报告"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def close(self):
|
||||
"""释放资源"""
|
||||
pass
|
||||
@@ -0,0 +1,19 @@
|
||||
"""
|
||||
可视化仓储接口 - 领域层
|
||||
定义活跃度可视化的抽象契约。
|
||||
"""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
|
||||
from ..models.data_models import ActivityVisualization
|
||||
|
||||
|
||||
class IActivityVisualizer(ABC):
|
||||
"""活跃度可视化接口 - 领域层抽象"""
|
||||
|
||||
@abstractmethod
|
||||
def generate_activity_visualization(
|
||||
self, messages: list[dict]
|
||||
) -> ActivityVisualization:
|
||||
"""从消息列表生成活跃度可视化数据"""
|
||||
pass
|
||||
@@ -0,0 +1,12 @@
|
||||
"""
|
||||
领域服务 - 分析业务逻辑服务
|
||||
|
||||
该模块导出所有封装核心业务逻辑的领域服务,
|
||||
用于分析群聊数据。这些服务是平台无关的。
|
||||
"""
|
||||
|
||||
from .incremental_merge_service import IncrementalMergeService
|
||||
|
||||
__all__ = [
|
||||
"IncrementalMergeService",
|
||||
]
|
||||
@@ -0,0 +1,142 @@
|
||||
"""
|
||||
分析领域服务 - 领域层
|
||||
负责用户维度的活跃度分析、发言习惯及活动模式识别。
|
||||
"""
|
||||
|
||||
from datetime import datetime
|
||||
from typing import TypedDict
|
||||
|
||||
from ..value_objects.unified_message import MessageContentType, UnifiedMessage
|
||||
|
||||
|
||||
class UserActivityStats(TypedDict):
|
||||
message_count: int
|
||||
char_count: int
|
||||
emoji_count: int
|
||||
nickname: str
|
||||
hours: dict[int, int]
|
||||
reply_count: int
|
||||
|
||||
|
||||
class AnalysisDomainService:
|
||||
"""分析领域服务 - 处理用户画像及行为分析"""
|
||||
|
||||
def analyze_user_activity(
|
||||
self,
|
||||
messages: list[UnifiedMessage],
|
||||
bot_self_ids: list[str] | None = None,
|
||||
) -> dict[str, UserActivityStats]:
|
||||
"""
|
||||
分析用户活跃度。
|
||||
|
||||
基于 UnifiedMessage 计算每个用户的发言数、字数、表情数等。
|
||||
"""
|
||||
user_stats: dict[str, UserActivityStats] = {}
|
||||
|
||||
bot_ids = set(bot_self_ids or [])
|
||||
|
||||
for msg in messages:
|
||||
user_id = msg.sender_id
|
||||
|
||||
# 跳过机器人自己的消息
|
||||
if user_id in bot_ids:
|
||||
continue
|
||||
|
||||
stats = user_stats.setdefault(
|
||||
user_id,
|
||||
{
|
||||
"message_count": 0,
|
||||
"char_count": 0,
|
||||
"emoji_count": 0,
|
||||
"nickname": "",
|
||||
"hours": {},
|
||||
"reply_count": 0,
|
||||
},
|
||||
)
|
||||
stats["message_count"] += 1
|
||||
stats["nickname"] = msg.sender_card or msg.sender_name
|
||||
|
||||
# 统计时间分布
|
||||
msg_time = datetime.fromtimestamp(msg.timestamp)
|
||||
hour = msg_time.hour
|
||||
stats["hours"][hour] = stats["hours"].get(hour, 0) + 1
|
||||
|
||||
# 统计内容
|
||||
for content in msg.contents:
|
||||
if content.type == MessageContentType.TEXT:
|
||||
stats["char_count"] += len(content.text or "")
|
||||
|
||||
elif content.type == MessageContentType.EMOJI:
|
||||
stats["emoji_count"] += 1
|
||||
|
||||
elif content.type == MessageContentType.IMAGE:
|
||||
# 与 GroupStatistics 口径保持一致
|
||||
if self._is_emoji_like_image(content.raw_data):
|
||||
stats["emoji_count"] += 1
|
||||
|
||||
elif content.type == MessageContentType.REPLY:
|
||||
stats["reply_count"] += 1
|
||||
|
||||
return user_stats
|
||||
|
||||
@staticmethod
|
||||
def _is_emoji_like_image(raw_data: object) -> bool:
|
||||
"""判断 IMAGE 段是否应按表情计数。"""
|
||||
if isinstance(raw_data, dict):
|
||||
sub_type = raw_data.get("sub_type")
|
||||
if sub_type is not None:
|
||||
return str(sub_type) == "1"
|
||||
summary = str(raw_data.get("summary", ""))
|
||||
return "动画表情" in summary or "表情" in summary
|
||||
|
||||
if raw_data is None:
|
||||
return False
|
||||
|
||||
text = str(raw_data)
|
||||
return "动画表情" in text or "表情" in text
|
||||
|
||||
def get_top_users(
|
||||
self, user_activity: dict[str, UserActivityStats], limit: int = 10
|
||||
) -> list[dict]:
|
||||
"""获取最活跃的用户列表"""
|
||||
users = []
|
||||
for user_id, stats in user_activity.items():
|
||||
users.append(
|
||||
{
|
||||
"user_id": user_id,
|
||||
"nickname": stats["nickname"],
|
||||
"message_count": stats["message_count"],
|
||||
"char_count": stats["char_count"],
|
||||
"emoji_count": stats["emoji_count"],
|
||||
"reply_count": stats["reply_count"],
|
||||
}
|
||||
)
|
||||
|
||||
# 按消息数量排序
|
||||
users.sort(key=lambda x: x["message_count"], reverse=True)
|
||||
return users[:limit]
|
||||
|
||||
def get_user_activity_pattern(
|
||||
self, user_activity: dict[str, UserActivityStats], user_id: str
|
||||
) -> dict:
|
||||
"""获取并识别指定用户的活动模式"""
|
||||
if user_id not in user_activity:
|
||||
return {}
|
||||
|
||||
stats = user_activity[user_id]
|
||||
hours = stats["hours"]
|
||||
|
||||
# 找出最活跃的时间段
|
||||
most_active_hour = max(hours.items(), key=lambda x: x[1])[0] if hours else 0
|
||||
|
||||
# 计算夜间活跃度 (0-6点)
|
||||
night_messages = sum(hours[h] for h in range(0, 6))
|
||||
night_ratio = (
|
||||
night_messages / stats["message_count"] if stats["message_count"] > 0 else 0
|
||||
)
|
||||
|
||||
return {
|
||||
"most_active_hour": most_active_hour,
|
||||
"night_ratio": night_ratio,
|
||||
"hourly_distribution": dict(hours),
|
||||
}
|
||||
@@ -0,0 +1,409 @@
|
||||
"""
|
||||
增量合并领域服务
|
||||
|
||||
负责将 IncrementalBatch 列表合并为 IncrementalState,
|
||||
以及将 IncrementalState 累积数据转换为现有实体类型,
|
||||
以便复用现有的报告生成器和分发器。
|
||||
|
||||
核心职责:
|
||||
- merge_batches: 将多个 IncrementalBatch 合并为一个 IncrementalState(滑动窗口聚合)
|
||||
- IncrementalState → GroupStatistics(含 ActivityVisualization、EmojiStatistics)
|
||||
- IncrementalState → list[SummaryTopic]
|
||||
- IncrementalState → list[GoldenQuote]
|
||||
"""
|
||||
|
||||
import time
|
||||
|
||||
from ...domain.entities.incremental_state import IncrementalBatch, IncrementalState
|
||||
from ...domain.models.data_models import (
|
||||
ActivityVisualization,
|
||||
EmojiStatistics,
|
||||
GoldenQuote,
|
||||
GroupStatistics,
|
||||
QualityDimension,
|
||||
QualityReview,
|
||||
SummaryTopic,
|
||||
TokenUsage,
|
||||
)
|
||||
from ...utils.logger import logger
|
||||
|
||||
|
||||
class IncrementalMergeService:
|
||||
"""
|
||||
增量合并服务
|
||||
|
||||
将滑动窗口内的多个批次数据合并为报告所需的数据结构,
|
||||
确保增量模式下生成的最终报告与传统单次分析报告格式完全一致。
|
||||
"""
|
||||
|
||||
def merge_batches(
|
||||
self,
|
||||
batches: list[IncrementalBatch],
|
||||
window_start: float,
|
||||
window_end: float,
|
||||
) -> IncrementalState:
|
||||
"""
|
||||
从批次列表合并构建 IncrementalState。
|
||||
|
||||
遍历所有批次,累加统计数据并对话题和金句执行去重,
|
||||
生成可用于报告的聚合视图。
|
||||
|
||||
Args:
|
||||
batches: 时间窗口内的批次列表(按时间升序)
|
||||
window_start: 窗口起始时间戳(epoch)
|
||||
window_end: 窗口结束时间戳(epoch)
|
||||
|
||||
Returns:
|
||||
IncrementalState: 合并后的聚合视图
|
||||
"""
|
||||
state = IncrementalState(
|
||||
group_id=batches[0].group_id if batches else "",
|
||||
window_start=window_start,
|
||||
window_end=window_end,
|
||||
total_analysis_count=len(batches),
|
||||
created_at=window_start,
|
||||
updated_at=time.time(),
|
||||
)
|
||||
|
||||
for batch in batches:
|
||||
# 累加消息和字符计数
|
||||
state.total_message_count += batch.messages_count
|
||||
state.total_character_count += batch.characters_count
|
||||
|
||||
# 合并每小时消息分布(按键累加)
|
||||
for hour_key, count in batch.hourly_msg_counts.items():
|
||||
hour_str = str(hour_key)
|
||||
state.hourly_message_counts[hour_str] = (
|
||||
state.hourly_message_counts.get(hour_str, 0) + count
|
||||
)
|
||||
|
||||
# 合并每小时字符分布
|
||||
for hour_key, count in batch.hourly_char_counts.items():
|
||||
hour_str = str(hour_key)
|
||||
state.hourly_character_counts[hour_str] = (
|
||||
state.hourly_character_counts.get(hour_str, 0) + count
|
||||
)
|
||||
|
||||
# 合并用户统计(按用户累加消息数、字符数等)
|
||||
for raw_user_id, stats in batch.user_stats.items():
|
||||
user_id = str(raw_user_id)
|
||||
if user_id not in state.user_activities:
|
||||
state.user_activities[user_id] = {
|
||||
"nickname": stats.get("nickname", stats.get("name", user_id)),
|
||||
"message_count": 0,
|
||||
"char_count": 0,
|
||||
"emoji_count": 0,
|
||||
"reply_count": 0,
|
||||
"hours": {},
|
||||
"last_message_time": 0,
|
||||
}
|
||||
existing = state.user_activities[user_id]
|
||||
existing["message_count"] += stats.get("message_count", 0)
|
||||
existing["char_count"] += stats.get("char_count", 0)
|
||||
existing["emoji_count"] += stats.get("emoji_count", 0)
|
||||
existing["reply_count"] += stats.get("reply_count", 0)
|
||||
|
||||
# 合并每小时统计
|
||||
# 兼容旧版本 (active_hours 是 list) 和新版本 (hours 是 dict)
|
||||
batch_hours = stats.get("hours", {})
|
||||
if isinstance(batch_hours, dict):
|
||||
# 现代 schema: hours 是 dict {hour: count}
|
||||
for h_str, h_count in batch_hours.items():
|
||||
h_int = int(h_str)
|
||||
existing["hours"][h_int] = (
|
||||
existing["hours"].get(h_int, 0) + h_count
|
||||
)
|
||||
else:
|
||||
# 兼容旧 schema: 只有 active_hours (list)
|
||||
active_hours = stats.get("active_hours", [])
|
||||
for h in active_hours:
|
||||
h_int = int(h)
|
||||
existing["hours"][h_int] = existing["hours"].get(h_int, 0) + 1
|
||||
|
||||
# 取最后消息时间的较大值
|
||||
batch_last = stats.get("last_message_time", 0)
|
||||
if batch_last > existing.get("last_message_time", 0):
|
||||
existing["last_message_time"] = batch_last
|
||||
|
||||
# 更新昵称(使用最新批次的有效昵称)
|
||||
nickname = stats.get("nickname", stats.get("name", ""))
|
||||
if nickname and str(nickname).strip():
|
||||
existing["nickname"] = nickname
|
||||
|
||||
# 合并表情统计(按键累加)
|
||||
for emoji_key, count in batch.emoji_stats.items():
|
||||
current_val = state.emoji_counts.get(emoji_key, 0)
|
||||
if isinstance(count, dict):
|
||||
# 如果是嵌套字典(如 face_details),则合并内部计数
|
||||
if not isinstance(current_val, dict):
|
||||
current_val = {}
|
||||
|
||||
for sub_key, sub_count in count.items():
|
||||
# 确保 current_val 是字典且 sub_count 是数字
|
||||
if isinstance(current_val, dict):
|
||||
current_val[sub_key] = (
|
||||
current_val.get(sub_key, 0) + sub_count
|
||||
)
|
||||
|
||||
state.emoji_counts[emoji_key] = current_val
|
||||
else:
|
||||
# 如果是数值,直接累加
|
||||
if isinstance(current_val, dict):
|
||||
# 异常情况:现有值是字典但新值是数字,通常不应发生,除非 schema 变更
|
||||
# 此时保留字典,忽略数字或记录错误,这里选择保留字典
|
||||
continue
|
||||
|
||||
state.emoji_counts[emoji_key] = current_val + count
|
||||
|
||||
# 合并话题(去重)
|
||||
for topic in batch.topics:
|
||||
if not IncrementalState.is_duplicate_topic(topic, state.topics):
|
||||
state.topics.append(topic)
|
||||
|
||||
# 合并金句(去重)
|
||||
for quote in batch.golden_quotes:
|
||||
if not IncrementalState.is_duplicate_quote(quote, state.golden_quotes):
|
||||
state.golden_quotes.append(quote)
|
||||
|
||||
# 累加 token 消耗
|
||||
for token_key in ("prompt_tokens", "completion_tokens", "total_tokens"):
|
||||
state.total_token_usage[token_key] = state.total_token_usage.get(
|
||||
token_key, 0
|
||||
) + batch.token_usage.get(token_key, 0)
|
||||
|
||||
# 合并参与者 ID(取并集)
|
||||
state.all_participant_ids.update(batch.participant_ids)
|
||||
|
||||
# 收集所有批次的质量锐评(用于最终汇总)
|
||||
if batch.chat_quality_review:
|
||||
state.all_quality_reviews.append(batch.chat_quality_review)
|
||||
|
||||
# 记录最后分析消息时间戳(取最大值)
|
||||
if batch.last_message_timestamp > state.last_analyzed_message_timestamp:
|
||||
state.last_analyzed_message_timestamp = batch.last_message_timestamp
|
||||
# 更新锐评为最新批次的 (如果没有汇总分析,则作为兜底)
|
||||
if batch.chat_quality_review:
|
||||
state.chat_quality_review = batch.chat_quality_review
|
||||
|
||||
logger.info(
|
||||
f"合并批次完成: 群={state.group_id}, "
|
||||
f"窗口={state.get_window_date_str()}, "
|
||||
f"批次数={len(batches)}, "
|
||||
f"总消息={state.total_message_count}, "
|
||||
f"话题={len(state.topics)}, 金句={len(state.golden_quotes)}"
|
||||
)
|
||||
|
||||
return state
|
||||
|
||||
def build_final_statistics(self, state: IncrementalState) -> GroupStatistics:
|
||||
"""
|
||||
从增量状态构建最终的群组统计数据。
|
||||
|
||||
将 IncrementalState 中的累积数据映射到 GroupStatistics,
|
||||
包含完整的 24 小时活跃度分布、表情统计和 token 消耗。
|
||||
|
||||
Args:
|
||||
state: 由 merge_batches 合并生成的增量分析状态
|
||||
|
||||
Returns:
|
||||
GroupStatistics: 与传统分析格式一致的统计数据
|
||||
"""
|
||||
# 构建 24 小时活跃度分布
|
||||
hourly_activity = {}
|
||||
for hour in range(24):
|
||||
hour_key = str(hour)
|
||||
hourly_activity[hour] = state.hourly_message_counts.get(hour_key, 0)
|
||||
|
||||
# 获取高峰时段
|
||||
peak_hours = state.get_peak_hours(3)
|
||||
|
||||
# 构建用户活跃排名
|
||||
user_ranking = state.get_user_activity_ranking(10)
|
||||
|
||||
# 构建活跃度可视化数据
|
||||
activity_visualization = ActivityVisualization(
|
||||
hourly_activity=hourly_activity,
|
||||
daily_activity={state.get_window_date_str(): state.total_message_count},
|
||||
user_activity_ranking=user_ranking,
|
||||
peak_hours=peak_hours,
|
||||
activity_heatmap_data={},
|
||||
)
|
||||
|
||||
# 构建表情统计
|
||||
emoji_statistics = self._build_emoji_statistics(state)
|
||||
|
||||
# 构建 token 消耗
|
||||
token_usage = TokenUsage(
|
||||
prompt_tokens=state.total_token_usage.get("prompt_tokens", 0),
|
||||
completion_tokens=state.total_token_usage.get("completion_tokens", 0),
|
||||
total_tokens=state.total_token_usage.get("total_tokens", 0),
|
||||
)
|
||||
|
||||
# 获取最活跃时段描述
|
||||
most_active_period = state.get_most_active_period()
|
||||
|
||||
# 转换聊天质量锐评 (如果有)
|
||||
chat_quality_review = None
|
||||
if state.chat_quality_review:
|
||||
review_dict = state.chat_quality_review
|
||||
dimensions_dict = review_dict.get("dimensions", [])
|
||||
dimensions = [
|
||||
QualityDimension(
|
||||
name=d.get("name", "未知"),
|
||||
percentage=float(d.get("percentage", 0)),
|
||||
comment=d.get("comment", ""),
|
||||
color=d.get("color", "#607d8b"),
|
||||
)
|
||||
for d in dimensions_dict
|
||||
]
|
||||
chat_quality_review = QualityReview(
|
||||
title=review_dict.get("title", "聊天质量锐评"),
|
||||
subtitle=review_dict.get("subtitle", "今天的群里发生了什么?"),
|
||||
dimensions=dimensions,
|
||||
summary=review_dict.get("summary", "今天也是充满活力的一天。"),
|
||||
)
|
||||
|
||||
statistics = GroupStatistics(
|
||||
message_count=state.total_message_count,
|
||||
total_characters=state.total_character_count,
|
||||
participant_count=len(state.all_participant_ids),
|
||||
most_active_period=most_active_period,
|
||||
golden_quotes=[], # 金句通过 build_quotes_for_report 单独构建
|
||||
emoji_count=emoji_statistics.total_emoji_count,
|
||||
emoji_statistics=emoji_statistics,
|
||||
activity_visualization=activity_visualization,
|
||||
token_usage=token_usage,
|
||||
chat_quality_review=chat_quality_review,
|
||||
)
|
||||
|
||||
logger.debug(
|
||||
f"从增量状态构建统计: "
|
||||
f"消息数={state.total_message_count}, "
|
||||
f"参与人数={len(state.all_participant_ids)}, "
|
||||
f"话题数={len(state.topics)}, "
|
||||
f"金句数={len(state.golden_quotes)}"
|
||||
)
|
||||
|
||||
return statistics
|
||||
|
||||
def build_topics_for_report(self, state: IncrementalState) -> list[SummaryTopic]:
|
||||
"""
|
||||
从增量状态构建报告用的话题列表。
|
||||
|
||||
将 IncrementalState 中累积的话题字典转换为 SummaryTopic 实例列表。
|
||||
|
||||
Args:
|
||||
state: 由 merge_batches 合并生成的增量分析状态
|
||||
|
||||
Returns:
|
||||
list[SummaryTopic]: 话题列表,格式与传统分析结果一致
|
||||
"""
|
||||
topics = []
|
||||
for topic_dict in state.topics:
|
||||
topic = SummaryTopic(
|
||||
topic=topic_dict.get("topic", "未知话题"),
|
||||
contributors=topic_dict.get("contributors", []),
|
||||
detail=topic_dict.get("detail", ""),
|
||||
contributor_ids=topic_dict.get("contributor_ids", []),
|
||||
)
|
||||
topics.append(topic)
|
||||
|
||||
logger.debug(f"从增量状态构建了 {len(topics)} 个话题")
|
||||
return topics
|
||||
|
||||
def build_quotes_for_report(self, state: IncrementalState) -> list[GoldenQuote]:
|
||||
"""
|
||||
从增量状态构建报告用的金句列表。
|
||||
|
||||
将 IncrementalState 中累积的金句字典转换为 GoldenQuote 实例列表。
|
||||
|
||||
Args:
|
||||
state: 由 merge_batches 合并生成的增量分析状态
|
||||
|
||||
Returns:
|
||||
list[GoldenQuote]: 金句列表,格式与传统分析结果一致
|
||||
"""
|
||||
quotes = []
|
||||
for quote_dict in state.golden_quotes:
|
||||
quote = GoldenQuote(
|
||||
content=quote_dict.get("content", ""),
|
||||
sender=quote_dict.get("sender", ""),
|
||||
reason=quote_dict.get("reason", ""),
|
||||
user_id=str(quote_dict.get("user_id", "")),
|
||||
)
|
||||
quotes.append(quote)
|
||||
|
||||
logger.debug(f"从增量状态构建了 {len(quotes)} 条金句")
|
||||
return quotes
|
||||
|
||||
def build_analysis_result(
|
||||
self,
|
||||
state: IncrementalState,
|
||||
user_titles: list | None = None,
|
||||
) -> dict:
|
||||
"""
|
||||
从增量状态构建完整的 analysis_result 字典。
|
||||
|
||||
该字典格式与 AnalysisApplicationService.execute_daily_analysis()
|
||||
返回的 analysis_result 完全一致,可直接传入 ReportDispatcher。
|
||||
|
||||
Args:
|
||||
state: 由 merge_batches 合并生成的增量分析状态
|
||||
user_titles: 用户称号列表(由最终报告时 LLM 分析生成)
|
||||
|
||||
Returns:
|
||||
dict: 包含 statistics、topics、user_titles、user_analysis 的结果字典
|
||||
"""
|
||||
statistics = self.build_final_statistics(state)
|
||||
topics = self.build_topics_for_report(state)
|
||||
golden_quotes = self.build_quotes_for_report(state)
|
||||
|
||||
# 将金句回填到 statistics 中(与传统流程一致)
|
||||
statistics.golden_quotes = golden_quotes
|
||||
|
||||
analysis_result = {
|
||||
"statistics": statistics,
|
||||
"topics": topics,
|
||||
"user_titles": user_titles or [],
|
||||
"user_analysis": state.user_activities,
|
||||
"chat_quality_review": statistics.chat_quality_review,
|
||||
}
|
||||
|
||||
logger.info(
|
||||
f"从增量状态构建完整分析结果: "
|
||||
f"群={state.group_id}, 窗口={state.get_window_date_str()}, "
|
||||
f"消息={state.total_message_count}, "
|
||||
f"话题={len(topics)}, "
|
||||
f"金句={len(golden_quotes)}, "
|
||||
f"批次={state.total_analysis_count}"
|
||||
)
|
||||
|
||||
return analysis_result
|
||||
|
||||
def _build_emoji_statistics(self, state: IncrementalState) -> EmojiStatistics:
|
||||
"""
|
||||
从增量状态构建表情统计。
|
||||
|
||||
将 IncrementalState 中的 emoji_counts 字典映射到 EmojiStatistics 字段。
|
||||
|
||||
Args:
|
||||
state: 增量分析状态
|
||||
|
||||
Returns:
|
||||
EmojiStatistics: 表情统计实例
|
||||
"""
|
||||
emoji_counts = state.emoji_counts
|
||||
|
||||
# 显式提取并检查类型,辅助 Pylance 类型推断
|
||||
face_details = emoji_counts.get("face_details")
|
||||
if not isinstance(face_details, dict):
|
||||
face_details = {}
|
||||
|
||||
return EmojiStatistics(
|
||||
face_count=emoji_counts.get("face_count", 0),
|
||||
mface_count=emoji_counts.get("mface_count", 0),
|
||||
bface_count=emoji_counts.get("bface_count", 0),
|
||||
sface_count=emoji_counts.get("sface_count", 0),
|
||||
other_emoji_count=emoji_counts.get("other_emoji_count", 0),
|
||||
face_details=face_details,
|
||||
)
|
||||
@@ -0,0 +1,103 @@
|
||||
"""
|
||||
消息清理服务 - 领域层
|
||||
负责过滤掉机器人消息、指令、技术性内容(如原始表情代码)及敏感内容。
|
||||
"""
|
||||
|
||||
import re
|
||||
from dataclasses import replace
|
||||
|
||||
from ..value_objects.unified_message import (
|
||||
MessageContent,
|
||||
MessageContentType,
|
||||
UnifiedMessage,
|
||||
)
|
||||
|
||||
# Discord 自定义表情正则 <:name:id> 或 <a:name:id>
|
||||
_DISCORD_CUSTOM_EMOJI_PATTERN = re.compile(r"<a?:.+?:\d+>")
|
||||
# 指令匹配正则:匹配以 / 开头,或者以 @某人 / 开头的消息
|
||||
_COMMAND_PATTERN = re.compile(r"^\s*(?:<@\d+>\s+)?/")
|
||||
|
||||
|
||||
class MessageCleanerService:
|
||||
"""消息清理服务"""
|
||||
|
||||
def clean_messages(
|
||||
self,
|
||||
messages: list[UnifiedMessage],
|
||||
bot_self_ids: list[str] = None,
|
||||
filter_commands: bool = True,
|
||||
) -> list[UnifiedMessage]:
|
||||
"""
|
||||
清理并过滤消息列表。
|
||||
|
||||
Args:
|
||||
messages: 原始统一格式消息列表
|
||||
bot_self_ids: 机器人自身的 ID 列表
|
||||
filter_commands: 是否过滤指令消息
|
||||
|
||||
Returns:
|
||||
清理后的消息列表
|
||||
"""
|
||||
bot_ids = set(bot_self_ids or [])
|
||||
cleaned_list = []
|
||||
|
||||
for msg in messages:
|
||||
# 1. 过滤机器人发送的消息
|
||||
if msg.sender_id in bot_ids:
|
||||
continue
|
||||
|
||||
# 2. 预检指令消息(首个内容块通常是文本)
|
||||
is_command = False
|
||||
first_text = msg.text_content
|
||||
if filter_commands and first_text and _COMMAND_PATTERN.match(first_text):
|
||||
is_command = True
|
||||
|
||||
if is_command:
|
||||
continue
|
||||
|
||||
# 3. 清理消息内容中的技术性噪音
|
||||
cleaned_contents = []
|
||||
has_meaningful_content = False
|
||||
|
||||
for content in msg.contents:
|
||||
if content.type == MessageContentType.TEXT:
|
||||
text = content.text or ""
|
||||
|
||||
# 移除 Discord 原始表情代码
|
||||
text = _DISCORD_CUSTOM_EMOJI_PATTERN.sub("", text)
|
||||
|
||||
# 移除 @mentions 文本 (e.g. <@123456>)
|
||||
text = re.sub(r"<@\d+>", "", text)
|
||||
|
||||
# 清理多余空格
|
||||
text = text.strip()
|
||||
|
||||
if text:
|
||||
cleaned_contents.append(
|
||||
MessageContent(type=MessageContentType.TEXT, text=text)
|
||||
)
|
||||
has_meaningful_content = True
|
||||
else:
|
||||
# 其他类型(图片、回复等)暂时保留,但由后续分析器决定是否使用
|
||||
cleaned_contents.append(content)
|
||||
if content.type != MessageContentType.REPLY:
|
||||
has_meaningful_content = True
|
||||
|
||||
# 4. 如果清理后仍有内容,则保留消息
|
||||
if has_meaningful_content:
|
||||
# 重新合成 text_content 用于 LLM 分析
|
||||
new_text_content = "".join(
|
||||
[
|
||||
c.text
|
||||
for c in cleaned_contents
|
||||
if c.type == MessageContentType.TEXT
|
||||
]
|
||||
).strip()
|
||||
|
||||
# 使用 replace 创建新实例(Frozen dataclass 必须如此)
|
||||
new_msg = replace(
|
||||
msg, contents=tuple(cleaned_contents), text_content=new_text_content
|
||||
)
|
||||
cleaned_list.append(new_msg)
|
||||
|
||||
return cleaned_list
|
||||
@@ -0,0 +1,130 @@
|
||||
"""
|
||||
统计领域服务 - 领域层
|
||||
负责核心统计逻辑的计算,不依赖于具体的平台或基础设施。
|
||||
"""
|
||||
|
||||
from collections import defaultdict
|
||||
from datetime import datetime
|
||||
|
||||
from ...infrastructure.visualization.activity_charts import ActivityVisualizer
|
||||
from ..models.data_models import EmojiStatistics, GroupStatistics, TokenUsage
|
||||
from ..repositories.visualization_repository import IActivityVisualizer
|
||||
from ..value_objects.unified_message import MessageContentType, UnifiedMessage
|
||||
|
||||
|
||||
class StatisticsService:
|
||||
"""统计服务 - 处理群聊数据的聚合统计"""
|
||||
|
||||
def __init__(self, activity_visualizer: IActivityVisualizer | None = None):
|
||||
if activity_visualizer is None:
|
||||
# Fallback: keep backward compatibility
|
||||
self.activity_visualizer: IActivityVisualizer = ActivityVisualizer()
|
||||
else:
|
||||
self.activity_visualizer = activity_visualizer
|
||||
|
||||
def calculate_group_statistics(
|
||||
self, messages: list[UnifiedMessage]
|
||||
) -> GroupStatistics:
|
||||
"""
|
||||
计算群组基础统计数据。
|
||||
|
||||
基于统一消息格式(UnifiedMessage)进行计算,确保跨平台一致性。
|
||||
"""
|
||||
total_chars = 0
|
||||
participants = set()
|
||||
hour_counts = defaultdict(int)
|
||||
emoji_statistics = EmojiStatistics()
|
||||
|
||||
for msg in messages:
|
||||
participants.add(msg.sender_id)
|
||||
|
||||
# 统计时间分布
|
||||
msg_time = datetime.fromtimestamp(msg.timestamp)
|
||||
hour_counts[msg_time.hour] += 1
|
||||
|
||||
# 处理消息内容
|
||||
for content in msg.contents:
|
||||
if content.type == MessageContentType.TEXT:
|
||||
total_chars += len(content.text or "")
|
||||
elif content.type == MessageContentType.EMOJI:
|
||||
emoji_statistics.face_count += 1
|
||||
# 尝试保留原始表情详情(如果适配器提供了)
|
||||
face_id = content.emoji_id or "unknown"
|
||||
emoji_statistics.face_details[f"emoji_{face_id}"] = (
|
||||
emoji_statistics.face_details.get(f"emoji_{face_id}", 0) + 1
|
||||
)
|
||||
elif content.type == MessageContentType.IMAGE:
|
||||
# 兼容识别“图片形态的表情”:
|
||||
# 1) 优先使用 onebot sub_type=1 信号
|
||||
# 2) 若无该字段,再回退到历史 summary 文本匹配
|
||||
if self._is_emoji_like_image(content.raw_data):
|
||||
emoji_statistics.mface_count += 1
|
||||
elif content.type in (
|
||||
MessageContentType.VOICE,
|
||||
MessageContentType.VIDEO,
|
||||
):
|
||||
# 其他非文本类型统计(可选)
|
||||
pass
|
||||
|
||||
# 找出最活跃时段
|
||||
most_active_hour = (
|
||||
max(hour_counts.items(), key=lambda x: x[1])[0] if hour_counts else 0
|
||||
)
|
||||
most_active_period = (
|
||||
f"{most_active_hour:02d}:00-{(most_active_hour + 1) % 24:02d}:00"
|
||||
)
|
||||
|
||||
# 生成活跃度可视化数据
|
||||
# 注意:ActivityVisualizer 可能需要迁移以支持 UnifiedMessage
|
||||
# 目前先转换回 dict 以保持兼容性,或者之后重构它
|
||||
raw_msgs = self._convert_to_legacy_dict(messages)
|
||||
activity_visualization = (
|
||||
self.activity_visualizer.generate_activity_visualization(raw_msgs)
|
||||
)
|
||||
|
||||
return GroupStatistics(
|
||||
message_count=len(messages),
|
||||
total_characters=total_chars,
|
||||
participant_count=len(participants),
|
||||
most_active_period=most_active_period,
|
||||
golden_quotes=[],
|
||||
emoji_count=emoji_statistics.total_emoji_count,
|
||||
emoji_statistics=emoji_statistics,
|
||||
activity_visualization=activity_visualization,
|
||||
token_usage=TokenUsage(),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _is_emoji_like_image(raw_data: object) -> bool:
|
||||
"""判断 IMAGE 段是否应按表情计数。"""
|
||||
if isinstance(raw_data, dict):
|
||||
sub_type = raw_data.get("sub_type")
|
||||
if sub_type is not None:
|
||||
return str(sub_type) == "1"
|
||||
summary = str(raw_data.get("summary", ""))
|
||||
return "动画表情" in summary or "表情" in summary
|
||||
|
||||
if raw_data is None:
|
||||
return False
|
||||
|
||||
text = str(raw_data)
|
||||
return "动画表情" in text or "表情" in text
|
||||
|
||||
def _convert_to_legacy_dict(self, messages: list[UnifiedMessage]) -> list[dict]:
|
||||
"""内部辅助:将 UnifiedMessage 转换为 Legacy Dict 格式,用于兼容可视化组件"""
|
||||
legacy_list = []
|
||||
for msg in messages:
|
||||
legacy_list.append(
|
||||
{
|
||||
"time": msg.timestamp,
|
||||
"sender": {
|
||||
"user_id": msg.sender_id,
|
||||
"nickname": msg.sender_name,
|
||||
"card": msg.sender_card or "",
|
||||
},
|
||||
"message": [
|
||||
{"type": "text", "data": {"text": msg.text_content or ""}}
|
||||
],
|
||||
}
|
||||
)
|
||||
return legacy_list
|
||||
@@ -0,0 +1,15 @@
|
||||
# 值对象
|
||||
from .platform_capabilities import PLATFORM_CAPABILITIES, PlatformCapabilities
|
||||
from .unified_group import UnifiedGroup, UnifiedMember
|
||||
from .unified_message import MessageContent, MessageContentType, UnifiedMessage
|
||||
|
||||
__all__ = [
|
||||
# 核心平台抽象
|
||||
"UnifiedMessage",
|
||||
"MessageContent",
|
||||
"MessageContentType",
|
||||
"PlatformCapabilities",
|
||||
"PLATFORM_CAPABILITIES",
|
||||
"UnifiedGroup",
|
||||
"UnifiedMember",
|
||||
]
|
||||
@@ -0,0 +1,305 @@
|
||||
"""
|
||||
平台能力值对象 - 运行时决策支持
|
||||
|
||||
每个平台适配器声明其能力,
|
||||
应用层根据能力决定操作。
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class PlatformCapabilities:
|
||||
"""
|
||||
值对象:平台能力描述
|
||||
|
||||
用于在运行时判断当前平台支持哪些具体操作,实现防御性编程和多平台兼容。
|
||||
|
||||
Attributes:
|
||||
platform_name (str): 平台标识(如 discord, onebot)
|
||||
platform_version (str): 版本号
|
||||
supports_message_history (bool): 是否支持拉取历史消息
|
||||
max_message_history_days (int): 最大历史穿透天数
|
||||
max_message_count (int): 单次拉取最大消息数
|
||||
supports_message_search (bool): 是否支持消息搜索(扩展用)
|
||||
supports_group_list (bool): 是否支持列出所有群组
|
||||
supports_group_info (bool): 是否支持获取群元数据
|
||||
supports_member_list (bool): 是否支持获取成员列表
|
||||
supports_member_info (bool): 是否支持获取单成员详情
|
||||
supports_text_message (bool): 是否能发送文本
|
||||
supports_image_message (bool): 是否能发送图片
|
||||
supports_file_message (bool): 是否能发送文件/PDF
|
||||
supports_forward_message (bool): 是否支持转发链(合并转发)
|
||||
supports_reply_message (bool): 是否支持回复引用
|
||||
max_text_length (int): 单条回复最大文本长度
|
||||
max_image_size_mb (float): 最大图片上传限制 (MB)
|
||||
supports_at_all (bool): 是否能 @全员
|
||||
supports_recall (bool): 是否支持撤回
|
||||
supports_edit (bool): 是否支持编辑已发消息
|
||||
supports_user_avatar (bool): 是否有用户头像 API
|
||||
supports_group_avatar (bool): 是否有群头像 API
|
||||
avatar_needs_api_call (bool): 获取头像是否需要额外异步请求
|
||||
avatar_sizes (tuple[int, ...]): 平台支持的头像尺寸像素值
|
||||
"""
|
||||
|
||||
# 平台标识
|
||||
platform_name: str
|
||||
platform_version: str = "unknown"
|
||||
|
||||
# 消息获取能力
|
||||
supports_message_history: bool = False
|
||||
max_message_history_days: int = 0
|
||||
max_message_count: int = 0
|
||||
supports_message_search: bool = False
|
||||
|
||||
# 群组信息能力
|
||||
supports_group_list: bool = False
|
||||
supports_group_info: bool = False
|
||||
supports_member_list: bool = False
|
||||
supports_member_info: bool = False
|
||||
|
||||
# 消息发送能力
|
||||
supports_text_message: bool = True
|
||||
supports_image_message: bool = False
|
||||
supports_file_message: bool = False
|
||||
supports_forward_message: bool = False
|
||||
supports_reply_message: bool = False
|
||||
max_text_length: int = 4096
|
||||
max_image_size_mb: float = 10.0
|
||||
|
||||
# 特殊能力
|
||||
supports_at_all: bool = False
|
||||
supports_recall: bool = False
|
||||
supports_edit: bool = False
|
||||
|
||||
# 头像能力
|
||||
supports_user_avatar: bool = True
|
||||
supports_group_avatar: bool = False
|
||||
avatar_needs_api_call: bool = False
|
||||
avatar_sizes: tuple[int, ...] = (100,)
|
||||
|
||||
# 检查方法
|
||||
def can_analyze(self) -> bool:
|
||||
"""
|
||||
判断是否具备进行群聊分析的核心能力。
|
||||
|
||||
Returns:
|
||||
bool: 核心能力齐全则返回 True
|
||||
"""
|
||||
return (
|
||||
self.supports_message_history
|
||||
and self.max_message_history_days > 0
|
||||
and self.max_message_count > 0
|
||||
)
|
||||
|
||||
def can_send_report(self, format: str = "image") -> bool:
|
||||
"""
|
||||
判断是否能以指定格式发送报告。
|
||||
|
||||
Args:
|
||||
format (str): 报告格式 ('text', 'image', 'pdf')
|
||||
|
||||
Returns:
|
||||
bool: 支持该格式则返回 True
|
||||
"""
|
||||
if format == "text":
|
||||
return self.supports_text_message
|
||||
elif format == "image":
|
||||
return self.supports_image_message
|
||||
elif format == "pdf":
|
||||
return self.supports_file_message
|
||||
return False
|
||||
|
||||
def get_effective_days(self, requested_days: int) -> int:
|
||||
"""
|
||||
获取实际生效的历史拉取天数。
|
||||
|
||||
Args:
|
||||
requested_days (int): 请求的天数
|
||||
|
||||
Returns:
|
||||
int: 平台受限后的实际天数
|
||||
"""
|
||||
return min(requested_days, self.max_message_history_days)
|
||||
|
||||
def get_effective_count(self, requested_count: int) -> int:
|
||||
"""
|
||||
获取实际生效的历史消息拉取条数。
|
||||
|
||||
Args:
|
||||
requested_count (int): 请求的消息条数
|
||||
|
||||
Returns:
|
||||
int: 平台受限后的实际条数
|
||||
"""
|
||||
return min(requested_count, self.max_message_count)
|
||||
|
||||
|
||||
# 预定义的平台能力
|
||||
# OneBot v11 (如 NapCat, LLOneBot 等)
|
||||
ONEBOT_V11_CAPABILITIES = PlatformCapabilities(
|
||||
platform_name="onebot",
|
||||
platform_version="v11",
|
||||
supports_message_history=True,
|
||||
max_message_history_days=7,
|
||||
max_message_count=10000,
|
||||
supports_group_list=True,
|
||||
supports_group_info=True,
|
||||
supports_member_list=True,
|
||||
supports_member_info=True,
|
||||
supports_text_message=True,
|
||||
supports_image_message=True,
|
||||
supports_file_message=True,
|
||||
supports_forward_message=True,
|
||||
supports_reply_message=True,
|
||||
max_text_length=4500,
|
||||
supports_at_all=True,
|
||||
supports_recall=True,
|
||||
supports_user_avatar=True,
|
||||
supports_group_avatar=True,
|
||||
avatar_needs_api_call=False,
|
||||
avatar_sizes=(40, 100, 140, 160, 640),
|
||||
)
|
||||
|
||||
# Telegram Bot API
|
||||
TELEGRAM_CAPABILITIES = PlatformCapabilities(
|
||||
platform_name="telegram",
|
||||
platform_version="bot_api_7.x",
|
||||
# 通过 PlatformMessageHistoryManager + 消息拦截器支持历史读取
|
||||
supports_message_history=True,
|
||||
max_message_history_days=7,
|
||||
max_message_count=1000,
|
||||
supports_group_list=False,
|
||||
supports_group_info=True,
|
||||
supports_member_list=True,
|
||||
supports_member_info=True,
|
||||
supports_text_message=True,
|
||||
supports_image_message=True,
|
||||
supports_file_message=True,
|
||||
supports_reply_message=True,
|
||||
max_text_length=4096,
|
||||
max_image_size_mb=50.0,
|
||||
supports_edit=True,
|
||||
supports_user_avatar=True,
|
||||
supports_group_avatar=True,
|
||||
avatar_needs_api_call=True,
|
||||
avatar_sizes=(160, 320, 640),
|
||||
)
|
||||
|
||||
# Discord API
|
||||
DISCORD_CAPABILITIES = PlatformCapabilities(
|
||||
platform_name="discord",
|
||||
platform_version="api_v10",
|
||||
supports_message_history=True,
|
||||
max_message_history_days=30,
|
||||
max_message_count=10000,
|
||||
supports_group_list=True,
|
||||
supports_group_info=True,
|
||||
supports_member_list=True,
|
||||
supports_text_message=True,
|
||||
supports_image_message=True,
|
||||
supports_file_message=True,
|
||||
supports_reply_message=True,
|
||||
max_text_length=2000,
|
||||
max_image_size_mb=8.0,
|
||||
supports_edit=True,
|
||||
supports_user_avatar=True,
|
||||
supports_group_avatar=True,
|
||||
avatar_needs_api_call=False,
|
||||
avatar_sizes=(16, 32, 64, 128, 256, 512, 1024, 2048, 4096),
|
||||
)
|
||||
|
||||
# Slack Web API
|
||||
SLACK_CAPABILITIES = PlatformCapabilities(
|
||||
platform_name="slack",
|
||||
platform_version="web_api",
|
||||
supports_message_history=True,
|
||||
max_message_history_days=90,
|
||||
max_message_count=1000,
|
||||
supports_group_list=True,
|
||||
supports_group_info=True,
|
||||
supports_member_list=True,
|
||||
supports_text_message=True,
|
||||
supports_image_message=True,
|
||||
supports_file_message=True,
|
||||
supports_reply_message=True,
|
||||
max_text_length=40000,
|
||||
supports_edit=True,
|
||||
supports_user_avatar=True,
|
||||
supports_group_avatar=False,
|
||||
avatar_needs_api_call=True,
|
||||
avatar_sizes=(24, 32, 48, 72, 192, 512, 1024),
|
||||
)
|
||||
|
||||
# Feishu/Lark Open Platform API
|
||||
LARK_CAPABILITIES = PlatformCapabilities(
|
||||
platform_name="lark",
|
||||
platform_version="open_api_v1",
|
||||
supports_message_history=True,
|
||||
max_message_history_days=7,
|
||||
max_message_count=1000,
|
||||
supports_group_list=False,
|
||||
supports_group_info=True,
|
||||
supports_member_list=True,
|
||||
supports_member_info=True,
|
||||
supports_text_message=True,
|
||||
supports_image_message=True,
|
||||
supports_file_message=True,
|
||||
supports_reply_message=True,
|
||||
max_text_length=30000,
|
||||
max_image_size_mb=10.0,
|
||||
supports_user_avatar=True,
|
||||
supports_group_avatar=True,
|
||||
avatar_needs_api_call=True,
|
||||
avatar_sizes=(72, 240, 640),
|
||||
)
|
||||
|
||||
# QQ Official Bot API. Message history is provided by the plugin's local
|
||||
# event archive because the public API does not expose group history queries.
|
||||
QQ_OFFICIAL_CAPABILITIES = PlatformCapabilities(
|
||||
platform_name="qq_official",
|
||||
platform_version="api_v2_local_history",
|
||||
supports_message_history=True,
|
||||
max_message_history_days=7,
|
||||
max_message_count=10000,
|
||||
supports_group_list=False,
|
||||
supports_group_info=False,
|
||||
supports_member_list=False,
|
||||
supports_member_info=False,
|
||||
supports_text_message=True,
|
||||
supports_image_message=True,
|
||||
supports_file_message=True,
|
||||
supports_forward_message=False,
|
||||
supports_reply_message=False,
|
||||
max_text_length=4000,
|
||||
max_image_size_mb=20.0,
|
||||
supports_user_avatar=True,
|
||||
supports_group_avatar=False,
|
||||
avatar_needs_api_call=False,
|
||||
avatar_sizes=(640,),
|
||||
)
|
||||
|
||||
# 能力查找表(映射平台标识到能力对象)
|
||||
PLATFORM_CAPABILITIES: dict[str, PlatformCapabilities] = {
|
||||
"aiocqhttp": ONEBOT_V11_CAPABILITIES,
|
||||
"onebot": ONEBOT_V11_CAPABILITIES,
|
||||
"telegram": TELEGRAM_CAPABILITIES,
|
||||
"discord": DISCORD_CAPABILITIES,
|
||||
"slack": SLACK_CAPABILITIES,
|
||||
"lark": LARK_CAPABILITIES,
|
||||
"qq_official": QQ_OFFICIAL_CAPABILITIES,
|
||||
"qq_official_webhook": QQ_OFFICIAL_CAPABILITIES,
|
||||
}
|
||||
|
||||
|
||||
def get_capabilities(platform_name: str) -> PlatformCapabilities | None:
|
||||
"""
|
||||
根据平台名称查找其支持的能力。
|
||||
|
||||
Args:
|
||||
platform_name (str): 平台名称
|
||||
|
||||
Returns:
|
||||
Optional[PlatformCapabilities]: 对应的能力对象或 None
|
||||
"""
|
||||
return PLATFORM_CAPABILITIES.get(platform_name.lower())
|
||||
@@ -0,0 +1,53 @@
|
||||
"""
|
||||
统一群组值对象 - 跨平台群组抽象
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class UnifiedMember:
|
||||
"""
|
||||
值对象:统一成员信息
|
||||
|
||||
Attributes:
|
||||
user_id (str): 用户唯一 ID
|
||||
nickname (str): 用户昵称
|
||||
card (str, optional): 群名片
|
||||
role (str): 角色(owner/admin/member)
|
||||
join_time (int, optional): 入群时间(秒级时间戳)
|
||||
avatar_url (str, optional): 头像网络链接
|
||||
avatar_data (str, optional): 头像 Base64 数据
|
||||
"""
|
||||
|
||||
user_id: str
|
||||
nickname: str
|
||||
card: str | None = None
|
||||
role: str = "member"
|
||||
join_time: int | None = None
|
||||
avatar_url: str | None = None
|
||||
avatar_data: str | None = None
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class UnifiedGroup:
|
||||
"""
|
||||
值对象:统一群组信息
|
||||
|
||||
Attributes:
|
||||
group_id (str): 群组唯一 ID
|
||||
group_name (str): 群组名称
|
||||
member_count (int): 成员数量
|
||||
owner_id (str, optional): 群主 ID
|
||||
create_time (int, optional): 创建时间
|
||||
description (str, optional): 群简介/公告
|
||||
platform (str): 来源平台
|
||||
"""
|
||||
|
||||
group_id: str
|
||||
group_name: str
|
||||
member_count: int = 0
|
||||
owner_id: str | None = None
|
||||
create_time: int | None = None
|
||||
description: str | None = None
|
||||
platform: str = "unknown"
|
||||
@@ -0,0 +1,177 @@
|
||||
"""
|
||||
统一消息值对象 - 跨平台核心抽象
|
||||
|
||||
所有平台消息都转换为此格式进行分析。
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime
|
||||
from enum import Enum
|
||||
from typing import Any
|
||||
|
||||
|
||||
class MessageContentType(Enum):
|
||||
"""
|
||||
枚举:消息内容类型
|
||||
|
||||
用于标识 MessageContent 的具体类型。
|
||||
"""
|
||||
|
||||
TEXT = "text"
|
||||
IMAGE = "image"
|
||||
FILE = "file"
|
||||
EMOJI = "emoji"
|
||||
REPLY = "reply"
|
||||
FORWARD = "forward"
|
||||
AT = "at"
|
||||
VOICE = "voice"
|
||||
VIDEO = "video"
|
||||
LOCATION = "location"
|
||||
UNKNOWN = "unknown"
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class MessageContent:
|
||||
"""
|
||||
值对象:消息内容段
|
||||
|
||||
表示消息链中的一个组成部分(如文本、图片、表情等)。
|
||||
该对象是不可变的,用于保证数据流的纯净。
|
||||
|
||||
Attributes:
|
||||
type (MessageContentType): 内容类型
|
||||
text (str): 文本内容(仅当类型为 TEXT 或包含文本描述时)
|
||||
url (str): 资源链接(图片、视频、文件等)
|
||||
emoji_id (str): 表情 ID
|
||||
emoji_name (str): 表情名称
|
||||
at_user_id (str): 被 @ 的用户 ID
|
||||
raw_data (Any): 平台原始数据,用于扩展
|
||||
"""
|
||||
|
||||
type: MessageContentType
|
||||
text: str = ""
|
||||
url: str = ""
|
||||
emoji_id: str = ""
|
||||
emoji_name: str = ""
|
||||
at_user_id: str = ""
|
||||
raw_data: Any = None
|
||||
|
||||
def is_text(self) -> bool:
|
||||
"""检查是否为文本内容。"""
|
||||
return self.type == MessageContentType.TEXT
|
||||
|
||||
def is_emoji(self) -> bool:
|
||||
"""检查是否为表情内容。"""
|
||||
return self.type == MessageContentType.EMOJI
|
||||
|
||||
@property
|
||||
def target_id(self) -> str:
|
||||
"""
|
||||
获取被 @ 的用户 ID(兼容旧代码)。
|
||||
|
||||
Alias for at_user_id.
|
||||
"""
|
||||
return self.at_user_id
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class UnifiedMessage:
|
||||
"""
|
||||
核心值对象:统一消息格式
|
||||
|
||||
跨平台抽象层,将不同平台的原始消息转换为统一格式进行分析。
|
||||
采用“只读”设计,确保分析逻辑的一致性。
|
||||
|
||||
Attributes:
|
||||
message_id (str): 消息唯一标识符
|
||||
sender_id (str): 发送者唯一 ID
|
||||
sender_name (str): 发送者昵称
|
||||
group_id (str): 群组/会话唯一 ID
|
||||
text_content (str): 经过清洗后的纯文本内容,主要用于 LLM 分析
|
||||
contents (tuple[MessageContent, ...]): 结构化消息链
|
||||
timestamp (int): Unix 时间戳(秒)
|
||||
platform (str): 来源平台名称(如 onebot, discord 等)
|
||||
reply_to_id (str, optional): 被回复的消息 ID
|
||||
sender_card (str, optional): 平台特定的群名片或特别备注
|
||||
"""
|
||||
|
||||
# 基础标识
|
||||
message_id: str
|
||||
sender_id: str
|
||||
sender_name: str
|
||||
group_id: str
|
||||
|
||||
# 消息内容
|
||||
text_content: str
|
||||
contents: tuple[MessageContent, ...] = field(default_factory=tuple)
|
||||
|
||||
# 时间信息
|
||||
timestamp: int = 0
|
||||
|
||||
# 平台信息
|
||||
platform: str = "unknown"
|
||||
|
||||
# 可选信息
|
||||
reply_to_id: str | None = None
|
||||
sender_card: str | None = None
|
||||
|
||||
# 分析辅助方法
|
||||
def has_text(self) -> bool:
|
||||
"""
|
||||
判断消息是否包含非空文本。
|
||||
|
||||
Returns:
|
||||
bool: 包含有效文本则返回 True
|
||||
"""
|
||||
return bool(self.text_content.strip())
|
||||
|
||||
def get_display_name(self) -> str:
|
||||
"""
|
||||
获取用户显示名称。
|
||||
优先级:群名片 > 昵称 > 用户 ID。
|
||||
|
||||
Returns:
|
||||
str: 格式化后的显示名称
|
||||
"""
|
||||
return self.sender_card or self.sender_name or self.sender_id
|
||||
|
||||
def get_emoji_count(self) -> int:
|
||||
"""
|
||||
计算消息链中包含的表情数量。
|
||||
|
||||
Returns:
|
||||
int: 表情总数
|
||||
"""
|
||||
return sum(1 for c in self.contents if c.is_emoji())
|
||||
|
||||
def get_text_length(self) -> int:
|
||||
"""
|
||||
获取文本内容的字符长度。
|
||||
|
||||
Returns:
|
||||
int: 字符数
|
||||
"""
|
||||
return len(self.text_content)
|
||||
|
||||
def get_datetime(self) -> datetime:
|
||||
"""
|
||||
将 Unix 时间戳转换为 datetime 对象。
|
||||
|
||||
Returns:
|
||||
datetime: 本地化后的时间对象
|
||||
"""
|
||||
return datetime.fromtimestamp(self.timestamp)
|
||||
|
||||
def to_analysis_format(self) -> str:
|
||||
"""
|
||||
转换为供 LLM 消费的分析格式。
|
||||
|
||||
Returns:
|
||||
str: 格式如 "[用户名]: 消息内容" 的字符串
|
||||
"""
|
||||
name = self.get_display_name()
|
||||
return f"[{name}]: {self.text_content}"
|
||||
|
||||
|
||||
# 类型别名
|
||||
MessageList = list[UnifiedMessage]
|
||||
@@ -0,0 +1 @@
|
||||
# 基础设施层
|
||||
@@ -0,0 +1,8 @@
|
||||
"""
|
||||
分析模块
|
||||
包含LLM分析功能
|
||||
"""
|
||||
|
||||
from .llm_analyzer import LLMAnalyzer
|
||||
|
||||
__all__ = ["LLMAnalyzer"]
|
||||
@@ -0,0 +1,11 @@
|
||||
"""
|
||||
分析器模块
|
||||
包含各种LLM分析功能的实现
|
||||
"""
|
||||
|
||||
from .base_analyzer import BaseAnalyzer
|
||||
from .golden_quote_analyzer import GoldenQuoteAnalyzer
|
||||
from .topic_analyzer import TopicAnalyzer
|
||||
from .user_title_analyzer import UserTitleAnalyzer
|
||||
|
||||
__all__ = ["BaseAnalyzer", "TopicAnalyzer", "UserTitleAnalyzer", "GoldenQuoteAnalyzer"]
|
||||
@@ -0,0 +1,659 @@
|
||||
"""
|
||||
基础分析器抽象类
|
||||
定义通用分析流程和接口
|
||||
"""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from collections.abc import Sized
|
||||
from typing import Generic, TypeVar
|
||||
|
||||
from ....domain.models.data_models import TokenUsage
|
||||
from ....shared.constants import PLUGIN_NAME
|
||||
from ....utils.logger import logger
|
||||
from ..utils.json_utils import parse_json_response
|
||||
from ..utils.llm_utils import (
|
||||
call_provider_with_retry,
|
||||
extract_response_text,
|
||||
extract_token_usage,
|
||||
get_provider_id_with_fallback,
|
||||
)
|
||||
from ..utils.structured_output_schema import JSONObject, build_response_format
|
||||
|
||||
TDataObject = TypeVar("TDataObject")
|
||||
TInputData = TypeVar("TInputData")
|
||||
|
||||
|
||||
class BaseAnalyzer(ABC, Generic[TDataObject, TInputData]):
|
||||
"""
|
||||
基础分析器抽象类
|
||||
定义所有分析器的通用接口 and 流程
|
||||
"""
|
||||
|
||||
def __init__(self, context, config_manager):
|
||||
"""
|
||||
初始化基础分析器
|
||||
|
||||
Args:
|
||||
context: AstrBot上下文对象
|
||||
config_manager: 配置管理器
|
||||
"""
|
||||
self.context = context
|
||||
self.config_manager = config_manager
|
||||
# 增量分析模式下的最大数量覆盖值,为 None 时使用配置默认值
|
||||
self._incremental_max_count: int | None = None
|
||||
|
||||
def get_provider_id_key(self) -> str | None:
|
||||
"""
|
||||
获取 Provider ID 配置键名
|
||||
子类可重写以指定特定的 provider,默认返回 None(使用主 LLM Provider)
|
||||
|
||||
Returns:
|
||||
Provider ID 配置键名,如 'topic_provider_id'
|
||||
"""
|
||||
return None
|
||||
|
||||
@abstractmethod
|
||||
def get_data_type(self) -> str:
|
||||
"""
|
||||
获取数据类型标识
|
||||
|
||||
Returns:
|
||||
数据类型字符串
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def get_max_count(self) -> int:
|
||||
"""
|
||||
获取最大提取数量
|
||||
|
||||
Returns:
|
||||
最大数量
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def build_prompt(self, data: TInputData) -> str:
|
||||
"""
|
||||
构建LLM提示词
|
||||
|
||||
Args:
|
||||
data: 输入数据
|
||||
|
||||
Returns:
|
||||
提示词字符串
|
||||
"""
|
||||
pass
|
||||
|
||||
def build_prompt_with_override(
|
||||
self, data: TInputData, prompt_override: str | None
|
||||
) -> str:
|
||||
"""构建提示词,默认忽略调用方覆盖模板。
|
||||
|
||||
Args:
|
||||
data: 分析器输入数据。
|
||||
prompt_override: 调用方指定的提示词模板。
|
||||
|
||||
Returns:
|
||||
可提交给 LLM 的提示词。
|
||||
"""
|
||||
del prompt_override
|
||||
return self.build_prompt(data)
|
||||
|
||||
@abstractmethod
|
||||
def extract_with_regex(self, result_text: str, max_count: int) -> list[dict]:
|
||||
"""
|
||||
使用正则表达式提取数据
|
||||
|
||||
Args:
|
||||
result_text: LLM响应文本
|
||||
max_count: 最大提取数量
|
||||
|
||||
Returns:
|
||||
提取到的数据列表
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def create_data_objects(self, data_list: list[dict]) -> list[TDataObject]:
|
||||
"""
|
||||
创建数据对象列表
|
||||
|
||||
Args:
|
||||
data_list: 原始数据列表
|
||||
|
||||
Returns:
|
||||
数据对象列表
|
||||
"""
|
||||
pass
|
||||
|
||||
def get_response_schema_name(self) -> str:
|
||||
return f"{self.get_data_type()}_output"
|
||||
|
||||
def get_response_schema(self) -> JSONObject | None:
|
||||
return None
|
||||
|
||||
def get_response_format(self) -> JSONObject | None:
|
||||
schema = self.get_response_schema()
|
||||
if not schema:
|
||||
return None
|
||||
return build_response_format(self.get_response_schema_name(), schema)
|
||||
|
||||
def get_schema_retry_max_attempts(self) -> int:
|
||||
"""
|
||||
schema 解析失败后的最大重试次数(不含首轮请求)。
|
||||
"""
|
||||
return 2
|
||||
|
||||
def get_schema_retry_temperatures(
|
||||
self, base_temperature: float | None
|
||||
) -> tuple[float, ...]:
|
||||
"""
|
||||
schema 解析失败后的温度重试序列(不含首轮请求)。
|
||||
采用动态降温,提高结构化稳定性。
|
||||
"""
|
||||
attempts = max(0, self.get_schema_retry_max_attempts())
|
||||
if attempts == 0:
|
||||
return ()
|
||||
|
||||
base = base_temperature if base_temperature is not None else 0.7
|
||||
first_retry = max(0.1, min(2.0, base * 0.5))
|
||||
|
||||
temperatures: list[float] = [round(first_retry, 2), 0.0]
|
||||
if attempts < len(temperatures):
|
||||
temperatures = temperatures[:attempts]
|
||||
|
||||
deduped: list[float] = []
|
||||
for temp in temperatures:
|
||||
if not deduped or deduped[-1] != temp:
|
||||
deduped.append(temp)
|
||||
return tuple(deduped)
|
||||
|
||||
async def _resolve_provider_temperature(
|
||||
self,
|
||||
provider_id_key: str | None,
|
||||
umo: str | None,
|
||||
provider_id: str | None = None,
|
||||
) -> float | None:
|
||||
"""
|
||||
尝试从当前将要调用的 Provider 配置中解析基础 temperature。
|
||||
"""
|
||||
# NoneBot 版:temperature 从插件配置读取;未配置则返回 None(交给请求方默认值)。
|
||||
try:
|
||||
return self.config_manager.get_llm_temperature()
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
def parse_structured_response(
|
||||
self, result_text: str
|
||||
) -> tuple[bool, list[dict] | None, str | None]:
|
||||
"""
|
||||
解析结构化响应(默认 JSON 数组解析)。
|
||||
子类可重写此方法定制对象解析逻辑。
|
||||
"""
|
||||
return parse_json_response(result_text, self.get_data_type())
|
||||
|
||||
def build_schema_retry_prompt(
|
||||
self,
|
||||
original_prompt: str,
|
||||
previous_output: str,
|
||||
parse_error: str | None,
|
||||
attempt_index: int,
|
||||
) -> str:
|
||||
"""
|
||||
构建结构化失败后的修复重试提示词。
|
||||
"""
|
||||
err_text = parse_error or "unknown_parse_error"
|
||||
return (
|
||||
f"{original_prompt}\n\n"
|
||||
"[STRUCTURED OUTPUT RETRY]\n"
|
||||
f"Attempt: {attempt_index}\n"
|
||||
"Your previous output did not satisfy the required strict JSON schema.\n"
|
||||
"Return ONLY valid JSON that strictly matches the schema. "
|
||||
"Do not include markdown, explanation, or extra text.\n"
|
||||
f"Parse error: {err_text}\n"
|
||||
"Previous invalid output:\n"
|
||||
f"{previous_output}"
|
||||
)
|
||||
|
||||
def _try_parse_with_fallback(
|
||||
self, result_text: str
|
||||
) -> tuple[bool, list[dict] | None, str | None]:
|
||||
"""
|
||||
先尝试结构化 JSON 解析(含修复逻辑),失败后立即尝试正则降级。
|
||||
"""
|
||||
success, parsed_data, error_msg = self.parse_structured_response(result_text)
|
||||
if success and parsed_data:
|
||||
validated_success, validated_data, validated_error = (
|
||||
self.validate_parsed_data(parsed_data)
|
||||
)
|
||||
if validated_success and validated_data:
|
||||
return True, validated_data, None
|
||||
error_msg = validated_error or error_msg
|
||||
|
||||
regex_data = self.extract_with_regex(result_text, self.get_max_count())
|
||||
if regex_data:
|
||||
validated_success, validated_data, validated_error = (
|
||||
self.validate_parsed_data(regex_data)
|
||||
)
|
||||
if validated_success and validated_data:
|
||||
logger.info(
|
||||
f"{self.get_data_type()}结构化解析失败后,正则降级提取成功,获得 {len(validated_data)} 条数据"
|
||||
)
|
||||
return True, validated_data, None
|
||||
error_msg = validated_error or error_msg
|
||||
|
||||
return False, None, error_msg
|
||||
|
||||
def validate_parsed_data(
|
||||
self, data_list: list[dict]
|
||||
) -> tuple[bool, list[dict] | None, str | None]:
|
||||
"""
|
||||
解析结果的本地二次校验(默认直接通过)。
|
||||
子类可重写为 Pydantic 校验。
|
||||
"""
|
||||
return True, data_list, None
|
||||
|
||||
def _save_debug_data(self, prompt: str, session_id: str):
|
||||
"""
|
||||
保存调试数据到文件
|
||||
|
||||
Args:
|
||||
prompt: 提示词内容
|
||||
session_id: 会话ID
|
||||
"""
|
||||
try:
|
||||
from ....utils.paths import get_data_dir
|
||||
|
||||
data_path = get_data_dir(PLUGIN_NAME) / "debug_data"
|
||||
|
||||
data_path.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
file_name = f"{session_id}_{self.get_data_type()}.txt"
|
||||
file_path = data_path / file_name
|
||||
|
||||
logger.info(f"正在保存调试数据到: {file_path}")
|
||||
with open(file_path, "w", encoding="utf-8") as f:
|
||||
f.write(prompt)
|
||||
|
||||
logger.info(f"已保存 {self.get_data_type()} 分析 Prompt 到 {file_path}")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"保存调试数据失败: {e}", exc_info=True)
|
||||
|
||||
def _apply_persona_reinforcement(
|
||||
self, prompt: str, system_prompt: str | None
|
||||
) -> str:
|
||||
"""
|
||||
核心的人格强化注入逻辑。采用首尾深度注入与指令交织策略。
|
||||
不仅强化输出口吻,更强调使用人格的逻辑视角进行分析过程。
|
||||
"""
|
||||
if not system_prompt or not system_prompt.strip():
|
||||
return prompt
|
||||
|
||||
logger.info(f"[{self.get_data_type()}分析] 已启用人格设定(深度强化模式)")
|
||||
|
||||
# 构造更具强制性的标识符
|
||||
persona_content = system_prompt.strip()
|
||||
|
||||
return (
|
||||
"【SYSTEM_CORE_IDENTITY_FIXED】\n"
|
||||
f"你现在的身份已由系统初始化为:\n{persona_content}\n\n"
|
||||
"--- MISSION_DIRECTIVE_START ---\n"
|
||||
"⚠️ 核心任务警告:你接下来的所有分析行为必须基于上述【身份设定】进行。\n"
|
||||
"这包括但不限于:你的思维切入点、对数据的敏感度、点评的犀利/温情程度、以及你对群聊氛围的感知逻辑。\n"
|
||||
f"请以该人格的思维方式去处理以下‘{self.get_data_type()}’分析任务:\n\n"
|
||||
f"{prompt}\n"
|
||||
"--- MISSION_DIRECTIVE_END ---\n\n"
|
||||
"【FINAL_IDENTITY_REINFORCEMENT】\n"
|
||||
f"1. 你不再是通用的 AI 助手,你是上述设定中的角色,我将在此处再次提醒你的身份:\n{persona_content}\n 正在观察并点评这些群聊数据。\n"
|
||||
f"2. 请务必使用该角色的第一人称视角 or 其独有的观察视角进行‘{self.get_data_type()}’输出。\n"
|
||||
"3. 你的分析成果必须体现该角色的性格色彩,禁止输出中立、客套、公式化的 AI 话术。\n"
|
||||
"4. ⚠️ 格式铁律:无论人格多么狂放,最终输出的内容必须严格遵守‘ MISSION_DIRECTIVE ’中所要求的纯 JSON 格式。除了 JSON 数据外,严禁输出任何 Markdown 标记或角色扮演的额外闲聊。"
|
||||
)
|
||||
|
||||
async def analyze(
|
||||
self,
|
||||
data: TInputData,
|
||||
umo: str | None = None,
|
||||
session_id: str | None = None,
|
||||
persona_id: str | None = None,
|
||||
prompt_override: str | None = None,
|
||||
) -> tuple[list[TDataObject], TokenUsage]:
|
||||
"""
|
||||
统一的分析流程
|
||||
|
||||
Args:
|
||||
data: 输入数据
|
||||
umo: 模型唯一标识符
|
||||
session_id: 会话ID (用于调试模式)
|
||||
persona_id: 显式指定的人格 ID,传入时优先于常规人格选择逻辑
|
||||
prompt_override: 调用方指定的提示词模板,供支持专属模板的分析器使用
|
||||
|
||||
Returns:
|
||||
(分析结果列表, Token使用统计)
|
||||
"""
|
||||
try:
|
||||
# 1. 构建提示词
|
||||
logger.debug(
|
||||
f"{self.get_data_type()}分析开始构建prompt,输入数据类型: {type(data)}"
|
||||
)
|
||||
data_length = len(data) if isinstance(data, Sized) else "N/A"
|
||||
logger.debug(f"{self.get_data_type()}分析输入数据长度: {data_length}")
|
||||
|
||||
prompt = self.build_prompt_with_override(data, prompt_override)
|
||||
logger.info(f"开始{self.get_data_type()}分析,构建提示词完成")
|
||||
logger.debug(
|
||||
f"{self.get_data_type()}分析prompt长度: {len(prompt) if prompt else 0}"
|
||||
)
|
||||
logger.debug(
|
||||
f"{self.get_data_type()}分析prompt前100字符: {prompt[:100] if prompt else 'None'}..."
|
||||
)
|
||||
|
||||
# 保存调试数据
|
||||
debug_mode = self.config_manager.get_debug_mode()
|
||||
if debug_mode and session_id and prompt:
|
||||
self._save_debug_data(prompt, session_id)
|
||||
elif debug_mode and not session_id:
|
||||
logger.warning("[Debug] Debug mode enabled but no session_id provided")
|
||||
|
||||
# 检查 prompt 是否为空
|
||||
if not prompt or not prompt.strip():
|
||||
logger.warning(
|
||||
f"{self.get_data_type()}分析: prompt 为空或只包含空白字符,跳过LLM调用"
|
||||
)
|
||||
return [], TokenUsage()
|
||||
|
||||
# 2. 调用LLM(使用配置的 provider)
|
||||
provider_id_key = self.get_provider_id_key()
|
||||
|
||||
# 只 resolve 一次 provider ID,同时传递给温度解析和 LLM 调用,避免重复日志
|
||||
resolved_provider_id = None
|
||||
if provider_id_key:
|
||||
resolved_provider_id = await get_provider_id_with_fallback(
|
||||
self.context, self.config_manager, provider_id_key, umo
|
||||
)
|
||||
|
||||
base_temperature = await self._resolve_provider_temperature(
|
||||
provider_id_key, umo, provider_id=resolved_provider_id
|
||||
)
|
||||
|
||||
# 获取人格设定
|
||||
system_prompt = await self._build_system_prompt(umo, persona_id)
|
||||
|
||||
# 应用人格强化注入
|
||||
prompt = self._apply_persona_reinforcement(prompt, system_prompt)
|
||||
|
||||
logger.info(f"[{self.get_data_type()}分析] 开始发起 LLM 请求, umo: {umo}")
|
||||
|
||||
# [Debug] 记录调试信息
|
||||
if debug_mode:
|
||||
logger.debug(
|
||||
f"[Debug] debug_mode={debug_mode}, umo={umo}, session_id={session_id}, prompt_len={len(prompt) if prompt else 0}"
|
||||
)
|
||||
|
||||
from ....shared.trace_context import TraceContext
|
||||
|
||||
trace = TraceContext.current()
|
||||
if trace:
|
||||
prompts_map = trace.metadata.setdefault("llm_prompts", {})
|
||||
prompts_map[self.get_data_type()] = {
|
||||
"prompt": prompt,
|
||||
"system_prompt": system_prompt,
|
||||
"provider_id": resolved_provider_id or "default",
|
||||
}
|
||||
|
||||
response = await call_provider_with_retry(
|
||||
self.context,
|
||||
self.config_manager,
|
||||
prompt=prompt,
|
||||
umo=umo,
|
||||
provider_id_key=provider_id_key,
|
||||
provider_id=resolved_provider_id,
|
||||
system_prompt=system_prompt,
|
||||
response_format=self.get_response_format(),
|
||||
observation_label=self.get_data_type(),
|
||||
)
|
||||
|
||||
if response is None:
|
||||
err_text = f"{self.get_data_type()}分析调用LLM失败: Provider 返回空响应或重试耗尽"
|
||||
logger.error(err_text)
|
||||
raise RuntimeError(err_text)
|
||||
|
||||
# 3. 提取token使用统计
|
||||
token_usage_dict = extract_token_usage(response)
|
||||
token_usage = TokenUsage(
|
||||
prompt_tokens=token_usage_dict["prompt_tokens"],
|
||||
completion_tokens=token_usage_dict["completion_tokens"],
|
||||
total_tokens=token_usage_dict["total_tokens"],
|
||||
)
|
||||
|
||||
# 4. 提取响应文本
|
||||
result_text = extract_response_text(response)
|
||||
logger.debug(f"{self.get_data_type()}分析原始响应: {result_text[:500]}...")
|
||||
|
||||
slot: dict | None = None
|
||||
if trace:
|
||||
prompts_map = trace.metadata.setdefault("llm_prompts", {})
|
||||
if isinstance(prompts_map, dict):
|
||||
slot = prompts_map.setdefault(self.get_data_type(), {})
|
||||
if isinstance(slot, dict):
|
||||
slot["prompt"] = prompt
|
||||
slot["initial_prompt"] = prompt
|
||||
slot["system_prompt"] = system_prompt
|
||||
slot["tokens"] = token_usage_dict["total_tokens"]
|
||||
slot["prompt_tokens"] = token_usage_dict["prompt_tokens"]
|
||||
slot["completion_tokens"] = token_usage_dict[
|
||||
"completion_tokens"
|
||||
]
|
||||
slot["completion"] = result_text
|
||||
slot["initial_completion"] = result_text
|
||||
|
||||
# 5. 尝试结构化解析 + 正则降级解析
|
||||
success, parsed_data, error_msg = self._try_parse_with_fallback(result_text)
|
||||
|
||||
# 5.1 仅在两种解析方式都失败时,进入 schema 修复重试(温度递减)
|
||||
if not success and self.get_response_format() is not None:
|
||||
temperatures = self.get_schema_retry_temperatures(base_temperature)
|
||||
for idx, temperature in enumerate(temperatures, start=1):
|
||||
retry_prompt = self.build_schema_retry_prompt(
|
||||
original_prompt=prompt,
|
||||
previous_output=result_text,
|
||||
parse_error=error_msg,
|
||||
attempt_index=idx,
|
||||
)
|
||||
logger.warning(
|
||||
f"{self.get_data_type()}结构化解析失败,触发 schema 修复重试 "
|
||||
f"(attempt={idx}, temperature={temperature:.1f})"
|
||||
)
|
||||
retry_response = await call_provider_with_retry(
|
||||
self.context,
|
||||
self.config_manager,
|
||||
prompt=retry_prompt,
|
||||
umo=umo,
|
||||
provider_id_key=provider_id_key,
|
||||
system_prompt=system_prompt,
|
||||
response_format=self.get_response_format(),
|
||||
extra_generate_kwargs={"temperature": temperature},
|
||||
observation_label=f"{self.get_data_type()}#schema_retry_{idx}",
|
||||
)
|
||||
if retry_response is None:
|
||||
continue
|
||||
|
||||
retry_result_text = extract_response_text(retry_response)
|
||||
if not retry_result_text:
|
||||
continue
|
||||
|
||||
result_text = retry_result_text
|
||||
retry_success, retry_parsed_data, retry_error_msg = (
|
||||
self._try_parse_with_fallback(retry_result_text)
|
||||
)
|
||||
if retry_success:
|
||||
success = True
|
||||
parsed_data = retry_parsed_data
|
||||
error_msg = None
|
||||
if trace and isinstance(slot, dict):
|
||||
slot["completion"] = retry_result_text
|
||||
slot["corrected_completion"] = retry_result_text
|
||||
slot["prompt"] = retry_prompt
|
||||
slot["corrected_prompt"] = retry_prompt
|
||||
slot["retry_count"] = idx
|
||||
retry_token_dict = extract_token_usage(retry_response)
|
||||
slot["tokens"] = (
|
||||
slot.get("tokens", 0) or 0
|
||||
) + retry_token_dict["total_tokens"]
|
||||
slot["completion_tokens"] = retry_token_dict[
|
||||
"completion_tokens"
|
||||
]
|
||||
break
|
||||
error_msg = retry_error_msg
|
||||
|
||||
if success and parsed_data:
|
||||
# JSON解析成功,创建数据对象
|
||||
data_objects = self.create_data_objects(parsed_data)
|
||||
logger.info(
|
||||
f"{self.get_data_type()}分析成功,解析到 {len(data_objects)} 条数据"
|
||||
)
|
||||
return data_objects, token_usage
|
||||
|
||||
# 6. 全部尝试失败
|
||||
err_text = f"{self.get_data_type()}分析失败: JSON解析与正则降级均未成功 ({error_msg or '未产出有效内容'})"
|
||||
logger.error(err_text)
|
||||
raise RuntimeError(err_text)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"{self.get_data_type()}分析失败: {e}", exc_info=True)
|
||||
raise
|
||||
|
||||
async def _build_system_prompt(
|
||||
self, umo: str | None, persona_id: str | None = None
|
||||
) -> str | None:
|
||||
"""
|
||||
构建带有会话人格的系统提示词,优先级如下:
|
||||
1. 调用方显式指定的人格
|
||||
2. 插件指定的全局人格 (若核心开关开启)
|
||||
3. 会话/对话选定的人格 (若开启了继承开关)
|
||||
4. 当前 UMO 的默认人格 (若开启了继承开关)
|
||||
|
||||
Args:
|
||||
umo: 用户模型对象标识,用于定位会话上下文
|
||||
persona_id: 调用方显式指定的人格 ID
|
||||
|
||||
Returns:
|
||||
最终生成的 System Prompt 字符串,若无则返回 None
|
||||
"""
|
||||
# 获取配置
|
||||
use_specific = self.config_manager.get_use_plugin_specific_persona()
|
||||
specific_id = self.config_manager.get_plugin_specific_persona_id()
|
||||
keep_original = self.config_manager.get_keep_original_persona()
|
||||
|
||||
# 获取 AstrBot 核心的人格管理器
|
||||
persona_mgr = getattr(self.context, "persona_manager", None)
|
||||
if persona_mgr is None:
|
||||
return None
|
||||
|
||||
persona_prompt = None
|
||||
|
||||
# 漫画角色等局部调用可显式指定人格,不修改插件全局人格配置。
|
||||
if persona_id:
|
||||
try:
|
||||
persona_obj = await persona_mgr.get_persona(persona_id)
|
||||
persona_prompt = (
|
||||
persona_obj.system_prompt
|
||||
if hasattr(persona_obj, "system_prompt")
|
||||
else None
|
||||
)
|
||||
if persona_prompt:
|
||||
logger.debug(f"已应用调用方指定人格: {persona_id}")
|
||||
except Exception as e:
|
||||
logger.warning(f"获取调用方指定人格失败 (ID: {persona_id}): {e}")
|
||||
|
||||
# --- 优先级 1: 插件指定的全局固定人格 ---
|
||||
# 适用于希望所有分析报告都呈现同一种风格的情况
|
||||
if not persona_prompt and use_specific and specific_id:
|
||||
try:
|
||||
persona_obj = await persona_mgr.get_persona(specific_id)
|
||||
persona_prompt = (
|
||||
persona_obj.system_prompt
|
||||
if hasattr(persona_obj, "system_prompt")
|
||||
else None
|
||||
)
|
||||
if persona_prompt:
|
||||
logger.debug(f"已应用插件指定的全局强制人格设定: {specific_id}")
|
||||
except Exception as e:
|
||||
logger.warning(f"获取插件指定人格失败 (ID: {specific_id}): {e}")
|
||||
|
||||
# --- 优先级 2: 继承当前会话/群聊的原始人格 ---
|
||||
# 只有在未开启“强制人格”且开启了“继承设定”时生效
|
||||
if not persona_prompt and keep_original and umo:
|
||||
try:
|
||||
# 2.1 尝试获取 SharedPreferences 中会话绑定的 Persona ID (通常是 /persona 命令设置的)
|
||||
from astrbot.api import sp
|
||||
|
||||
session_service_config = await sp.get_async(
|
||||
scope="umo",
|
||||
scope_id=str(umo),
|
||||
key="session_service_config",
|
||||
default={},
|
||||
)
|
||||
persona_id = (
|
||||
session_service_config.get("persona_id")
|
||||
if session_service_config
|
||||
else None
|
||||
)
|
||||
|
||||
if persona_id and persona_id != "[%None]":
|
||||
persona_obj = await persona_mgr.get_persona(persona_id)
|
||||
persona_prompt = (
|
||||
persona_obj.system_prompt
|
||||
if hasattr(persona_obj, "system_prompt")
|
||||
else None
|
||||
)
|
||||
if persona_prompt:
|
||||
logger.debug(f"继承到会话选定人格: {persona_id}")
|
||||
|
||||
# 2.2 若无会话绑定,尝试获取当前对话(Dialogue)级别的人格
|
||||
if not persona_prompt:
|
||||
conv_mgr = getattr(self.context, "conversation_manager", None)
|
||||
if conv_mgr:
|
||||
curr_conv_id = await conv_mgr.get_curr_conversation_id(umo)
|
||||
if curr_conv_id:
|
||||
conv_obj = await conv_mgr.get_conversation(
|
||||
umo, curr_conv_id
|
||||
)
|
||||
if (
|
||||
conv_obj
|
||||
and conv_obj.persona_id
|
||||
and conv_obj.persona_id != "[%None]"
|
||||
):
|
||||
persona_obj = await persona_mgr.get_persona(
|
||||
conv_obj.persona_id
|
||||
)
|
||||
persona_prompt = (
|
||||
persona_obj.system_prompt
|
||||
if hasattr(persona_obj, "system_prompt")
|
||||
else None
|
||||
)
|
||||
if persona_prompt:
|
||||
logger.debug(
|
||||
f"继承到对话(Dialogue)设定人格: {conv_obj.persona_id}"
|
||||
)
|
||||
|
||||
# 2.3 若仍无结果,尝试获取 UMO 设定的默认人格
|
||||
if not persona_prompt:
|
||||
personality = await persona_mgr.get_default_persona_v3(umo)
|
||||
if isinstance(personality, dict):
|
||||
persona_prompt = personality.get("prompt")
|
||||
else:
|
||||
persona_prompt = getattr(personality, "prompt", None)
|
||||
if persona_prompt:
|
||||
logger.debug("继承到 UMO 默认人格设定")
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"分析人格回溯识别失败 (umo: {umo}): {e}")
|
||||
|
||||
# 检查生成结果
|
||||
if not isinstance(persona_prompt, str) or not persona_prompt.strip():
|
||||
return None
|
||||
|
||||
return persona_prompt.strip()
|
||||
@@ -0,0 +1,614 @@
|
||||
"""
|
||||
聊天质量分析模块
|
||||
专门处理群聊质量锐评分析
|
||||
"""
|
||||
|
||||
from datetime import datetime
|
||||
|
||||
from ....domain.models.data_models import QualityDimension, QualityReview, TokenUsage
|
||||
from ....shared.trace_context import TraceContext
|
||||
from ....utils.logger import logger
|
||||
from ...utils.template_utils import render_template
|
||||
from ..utils import InfoUtils
|
||||
from ..utils.json_utils import extract_quality_with_regex, parse_json_object_response
|
||||
from ..utils.llm_utils import (
|
||||
call_provider_with_retry,
|
||||
extract_response_text,
|
||||
extract_token_usage,
|
||||
get_provider_id_with_fallback,
|
||||
)
|
||||
from ..utils.response_validation import validate_quality_review_item
|
||||
from ..utils.structured_output_schema import JSONObject, build_chat_quality_schema
|
||||
from .base_analyzer import BaseAnalyzer
|
||||
|
||||
|
||||
class ChatQualityAnalyzer(BaseAnalyzer[QualityReview, list[dict]]):
|
||||
"""
|
||||
聊天质量分析器
|
||||
专门处理群聊质量的锐评和多维度分析
|
||||
|
||||
注意:由于聊天质量分析返回的是 JSON 对象而非数组,
|
||||
此分析器重写了 analyze() 方法,使用 parse_json_object_response 解析,
|
||||
并以 extract_quality_with_regex 作为正则降级方案。
|
||||
"""
|
||||
|
||||
def get_provider_id_key(self) -> str:
|
||||
"""获取 Provider ID 配置键名"""
|
||||
return "quality_provider_id"
|
||||
|
||||
def get_data_type(self) -> str:
|
||||
"""获取数据类型标识"""
|
||||
return "聊天质量"
|
||||
|
||||
def get_max_count(self) -> int:
|
||||
"""获取最大维度数量"""
|
||||
return 8
|
||||
|
||||
def get_response_schema_name(self) -> str:
|
||||
return "daily_chat_quality_review"
|
||||
|
||||
def get_response_schema(self) -> JSONObject:
|
||||
return build_chat_quality_schema(self.get_max_count())
|
||||
|
||||
def build_prompt(self, data: list[dict]) -> str:
|
||||
"""
|
||||
构建聊天质量分析提示词
|
||||
"""
|
||||
if not data:
|
||||
return ""
|
||||
|
||||
# 提取文本消息
|
||||
text_messages = []
|
||||
for msg in data:
|
||||
if not isinstance(msg, dict):
|
||||
continue
|
||||
|
||||
sender = msg.get("sender", {})
|
||||
user_id = str(sender.get("user_id", ""))
|
||||
bot_self_ids = self.config_manager.get_bot_self_ids()
|
||||
if bot_self_ids and user_id in [str(uid) for uid in bot_self_ids]:
|
||||
continue
|
||||
|
||||
nickname = InfoUtils.get_user_nickname(self.config_manager, sender)
|
||||
msg_time = datetime.fromtimestamp(msg.get("time", 0)).strftime("%H:%M")
|
||||
message_list = msg.get("message", [])
|
||||
|
||||
text_parts = []
|
||||
for content in message_list:
|
||||
if content.get("type") == "text":
|
||||
text = content.get("data", {}).get("text", "").strip()
|
||||
if text:
|
||||
text_parts.append(text)
|
||||
|
||||
combined_text = "".join(text_parts).strip()
|
||||
if combined_text and not combined_text.startswith("/"):
|
||||
text_messages.append(f"[{msg_time}] [{nickname}]: {combined_text}")
|
||||
|
||||
messages_text = "\n".join(text_messages[:1000])
|
||||
|
||||
prompt_template = self.config_manager.get_quality_analysis_prompt()
|
||||
|
||||
if prompt_template:
|
||||
return render_template(prompt_template, messages_text=messages_text)
|
||||
|
||||
prompt_template = """请分析以下群聊记录,输出一份"聊天质量锐评"。
|
||||
|
||||
## 任务目标:
|
||||
1. **维度划分**:将聊天内容划分为 3-6 个【高层级、抽象、泛化】的维度(例如:就业焦虑、生涯规划、技术方案研究、情感树洞、无意义水群等)。
|
||||
2. **严禁在维度名称(name)中出现任何具体的群聊人物名、项目名、具体的报错内容或细碎的事件点。标题必须保持高度抽象且字数简练(2-6个字)。**
|
||||
3. 为每个维度计算一个大致的百分比占位(总和小于等于 100%)。
|
||||
4. **点评内容**:为每个维度写一句犀利、幽默、毒舌或温情的点评。具体的吐槽内容、具体的细节事件描述请放在这里。
|
||||
5. **全群表现**:给出一句总结性的评价,作为总结标题对应的“金句”。
|
||||
6. **主题设定**:设定一个本次报告的主题标题和副标题。
|
||||
|
||||
## 点评风格指南:
|
||||
- 语言要接地气,多用互联网黑话。吐槽要精准,避重就轻。
|
||||
- **只有维度名称(name)需要抽象,点评(comment)和总结(summary)可以非常具体和生动。**
|
||||
|
||||
## 返回格式要求:
|
||||
必须以纯 JSON 格式返回,不得包含任何 Markdown 格式。
|
||||
|
||||
```json
|
||||
{{
|
||||
"title": "今日群聊主题",
|
||||
"subtitle": "副标题",
|
||||
"dimensions": [
|
||||
{{
|
||||
"name": "抽象维度名",
|
||||
"percentage": 比例,
|
||||
"comment": "维度的毒舌点评"
|
||||
}}
|
||||
],
|
||||
"summary": "一句总结性的金句"
|
||||
}}
|
||||
```
|
||||
|
||||
群聊记录:
|
||||
${messages_text}
|
||||
"""
|
||||
|
||||
return render_template(prompt_template, messages_text=messages_text)
|
||||
|
||||
def extract_with_regex(self, result_text: str, max_count: int) -> list[dict]:
|
||||
"""
|
||||
使用正则表达式提取质量分析数据(BaseAnalyzer 要求的接口)
|
||||
|
||||
注意: 此方法供 BaseAnalyzer.analyze() 的降级流程使用,
|
||||
但由于聊天质量分析重写了 analyze(),实际由 analyze_quality() 中调用
|
||||
extract_quality_with_regex 实现。
|
||||
"""
|
||||
return []
|
||||
|
||||
def create_data_objects(self, data_list: list[dict]) -> list[QualityReview]:
|
||||
"""
|
||||
满足 BaseAnalyzer 抽象要求。
|
||||
聊天质量分析的数据对象创建在 analyze_quality 中完成。
|
||||
"""
|
||||
return []
|
||||
|
||||
def _build_review_from_dict(self, data: dict) -> QualityReview:
|
||||
"""
|
||||
从解析后的字典构建 QualityReview 对象
|
||||
|
||||
Args:
|
||||
data: 解析后的 JSON 对象字典
|
||||
|
||||
Returns:
|
||||
QualityReview 数据对象
|
||||
"""
|
||||
# 控制维度占比总和不超过100%
|
||||
total_percentage = sum(
|
||||
max(0.0, min(100.0, float(d.get("percentage", 0))))
|
||||
for d in data.get("dimensions", [])
|
||||
)
|
||||
|
||||
factor = 1.0
|
||||
if total_percentage > 100:
|
||||
factor = 100.0 / total_percentage
|
||||
|
||||
dimensions = []
|
||||
for d in data.get("dimensions", []):
|
||||
raw_p = float(d.get("percentage", 0))
|
||||
|
||||
final_p = round(max(0.0, min(100.0, raw_p)) * factor, 1)
|
||||
|
||||
dimensions.append(
|
||||
QualityDimension(
|
||||
name=d.get("name", "未知"),
|
||||
percentage=final_p,
|
||||
comment=d.get("comment", ""),
|
||||
)
|
||||
)
|
||||
|
||||
# 自动分配颜色
|
||||
colors = [
|
||||
"#607d8b",
|
||||
"#2196f3",
|
||||
"#f44336",
|
||||
"#e91e63",
|
||||
"#ff9800",
|
||||
"#4caf50",
|
||||
"#009688",
|
||||
"#9c27b0",
|
||||
]
|
||||
for i, d in enumerate(dimensions):
|
||||
d.color = colors[i % len(colors)]
|
||||
|
||||
return QualityReview(
|
||||
title=data.get("title", "聊天质量锐评"),
|
||||
subtitle=data.get("subtitle", "今天的群里发生了什么?"),
|
||||
dimensions=dimensions,
|
||||
summary=data.get("summary", "今天也是充满活力的一天。"),
|
||||
)
|
||||
|
||||
def _validate_review_payload(
|
||||
self, data: dict
|
||||
) -> tuple[bool, dict | None, str | None]:
|
||||
return validate_quality_review_item(data)
|
||||
|
||||
async def _retry_parse_quality_object(
|
||||
self,
|
||||
*,
|
||||
original_prompt: str,
|
||||
previous_output: str,
|
||||
parse_error: str | None,
|
||||
umo: str | None,
|
||||
system_prompt: str | None,
|
||||
base_temperature: float | None,
|
||||
) -> dict | None:
|
||||
response_format = self.get_response_format()
|
||||
if response_format is None:
|
||||
return None
|
||||
|
||||
for idx, temperature in enumerate(
|
||||
self.get_schema_retry_temperatures(base_temperature), start=1
|
||||
):
|
||||
retry_prompt = self.build_schema_retry_prompt(
|
||||
original_prompt=original_prompt,
|
||||
previous_output=previous_output,
|
||||
parse_error=parse_error,
|
||||
attempt_index=idx,
|
||||
)
|
||||
logger.warning(
|
||||
f"聊天质量结构化解析失败,触发 schema 修复重试 "
|
||||
f"(attempt={idx}, temperature={temperature:.1f})"
|
||||
)
|
||||
retry_response = await call_provider_with_retry(
|
||||
self.context,
|
||||
self.config_manager,
|
||||
prompt=retry_prompt,
|
||||
umo=umo,
|
||||
provider_id_key=self.get_provider_id_key(),
|
||||
system_prompt=system_prompt,
|
||||
response_format=response_format,
|
||||
extra_generate_kwargs={"temperature": temperature},
|
||||
observation_label=f"{self.get_data_type()}#schema_retry_{idx}",
|
||||
)
|
||||
if retry_response is None:
|
||||
continue
|
||||
|
||||
retry_text = extract_response_text(retry_response)
|
||||
if not retry_text:
|
||||
continue
|
||||
|
||||
from ....shared.trace_context import TraceContext
|
||||
|
||||
trace = TraceContext.current()
|
||||
|
||||
retry_success, retry_parsed_data, _ = parse_json_object_response(
|
||||
retry_text, self.get_data_type()
|
||||
)
|
||||
if retry_success and retry_parsed_data:
|
||||
valid, normalized, _ = self._validate_review_payload(retry_parsed_data)
|
||||
if valid and normalized:
|
||||
if trace:
|
||||
slot = trace.metadata.get("llm_prompts", {}).get(
|
||||
self.get_data_type()
|
||||
)
|
||||
if isinstance(slot, dict):
|
||||
slot["completion"] = retry_text
|
||||
slot["corrected_completion"] = retry_text
|
||||
slot["prompt"] = retry_prompt
|
||||
slot["corrected_prompt"] = retry_prompt
|
||||
slot["retry_count"] = idx
|
||||
return normalized
|
||||
|
||||
retry_regex_data = extract_quality_with_regex(retry_text)
|
||||
if retry_regex_data:
|
||||
valid, normalized, _ = self._validate_review_payload(retry_regex_data)
|
||||
if valid and normalized:
|
||||
if trace:
|
||||
slot = trace.metadata.get("llm_prompts", {}).get(
|
||||
self.get_data_type()
|
||||
)
|
||||
if isinstance(slot, dict):
|
||||
slot["completion"] = retry_text
|
||||
slot["corrected_completion"] = retry_text
|
||||
slot["prompt"] = retry_prompt
|
||||
slot["corrected_prompt"] = retry_prompt
|
||||
slot["retry_count"] = idx
|
||||
return normalized
|
||||
|
||||
return None
|
||||
|
||||
async def summarize_batch_reviews(
|
||||
self,
|
||||
batch_reviews: list[dict],
|
||||
umo: str | None = None,
|
||||
session_id: str | None = None,
|
||||
) -> tuple[QualityReview | None, TokenUsage]:
|
||||
"""
|
||||
汇总多个增量批次的质量报告,生成最终的每日全天总评。
|
||||
"""
|
||||
if not batch_reviews:
|
||||
return None, TokenUsage()
|
||||
|
||||
if len(batch_reviews) == 1:
|
||||
return self._build_review_from_dict(batch_reviews[0]), TokenUsage()
|
||||
|
||||
try:
|
||||
# 构建汇总用的提示词
|
||||
reviews_text = ""
|
||||
for i, rev in enumerate(batch_reviews):
|
||||
title = rev.get("title", "未命名")
|
||||
summary = rev.get("summary", "")
|
||||
dims = ", ".join(
|
||||
[
|
||||
f"{d.get('name')}({d.get('percentage')}%)"
|
||||
for d in rev.get("dimensions", [])
|
||||
]
|
||||
)
|
||||
reviews_text += f"\n批次 {i + 1} [{title}]:\n- 维度表现: {dims}\n- 核心摘要: {summary}\n"
|
||||
|
||||
# 获取配置中的汇总提示词模板,如果没有则使用默认模板
|
||||
prompt_template = (
|
||||
self.config_manager.get_quality_summary_prompt()
|
||||
or """你现在有一份今天全天分散时间段的多个“增量批次点评笔记”。
|
||||
你的任务是将这些分散的笔记汇总成一份最终的“全天聊天质量终极锐评”。
|
||||
|
||||
## 任务目标:
|
||||
1. **全局抽象维度**:根据各批次的维度表现,平衡权重,提取出 3-6 个覆盖全天的【核心、上层抽象】课题维度(如:职场/行业风向、技术架构演进、社畜心理博弈等)。
|
||||
2. **严禁在维度名称(name)中出现具体的批次细节。标题必须代表全天的某种趋势。**
|
||||
3. **百分比融合**:根据全天笔记的频率和强度,给出一个代表全天整体分布的比例(总和不超过100%)。
|
||||
4. **终极点评**:为每个汇总维度写出一句升华后的全天总结性点评。可以融合具体批次中的有趣槽点。
|
||||
5. **终极总结**:拟定全天的大型主题标题、副标题,并给出一句霸气的全天表现总结。
|
||||
|
||||
## 风格要求:
|
||||
- 只有维度名称(name)需要高度概括抽象。
|
||||
- 点评(comment)和总结(summary)请尽量生动、具体,要把一整天的梗串联起来。
|
||||
|
||||
## 返回格式要求:
|
||||
必须以纯 JSON 格式返回,不得包含任何 Markdown 格式。
|
||||
|
||||
```json
|
||||
{{
|
||||
"title": "今日群聊主题",
|
||||
"subtitle": "副标题",
|
||||
"dimensions": [
|
||||
{{
|
||||
"name": "抽象大类标题",
|
||||
"percentage": 比例,
|
||||
"comment": "维度的全天锐评"
|
||||
}}
|
||||
],
|
||||
"summary": "全天总结金句"
|
||||
}}
|
||||
```
|
||||
"""
|
||||
)
|
||||
prompt = render_template(prompt_template, reviews_text=reviews_text)
|
||||
|
||||
# 调用 LLM 进行汇总
|
||||
system_prompt = await self._build_system_prompt(umo)
|
||||
base_temperature = await self._resolve_provider_temperature(
|
||||
self.get_provider_id_key(), umo
|
||||
)
|
||||
|
||||
# 应用人设强化注入
|
||||
prompt = self._apply_persona_reinforcement(prompt, system_prompt)
|
||||
|
||||
response = await call_provider_with_retry(
|
||||
self.context,
|
||||
self.config_manager,
|
||||
prompt=prompt,
|
||||
umo=umo,
|
||||
provider_id_key=self.get_provider_id_key(),
|
||||
system_prompt=system_prompt,
|
||||
response_format=self.get_response_format(),
|
||||
observation_label=f"{self.get_data_type()}#batch_summary",
|
||||
)
|
||||
|
||||
if response is None:
|
||||
return None, TokenUsage()
|
||||
|
||||
token_usage_dict = extract_token_usage(response)
|
||||
usage = TokenUsage(
|
||||
prompt_tokens=token_usage_dict["prompt_tokens"],
|
||||
completion_tokens=token_usage_dict["completion_tokens"],
|
||||
total_tokens=token_usage_dict["total_tokens"],
|
||||
)
|
||||
|
||||
result_text = extract_response_text(response)
|
||||
trace = TraceContext.current()
|
||||
if trace:
|
||||
slot = trace.metadata.setdefault("llm_prompts", {}).setdefault(
|
||||
self.get_data_type(), {}
|
||||
)
|
||||
slot["prompt"] = prompt
|
||||
slot["initial_prompt"] = prompt
|
||||
slot["system_prompt"] = system_prompt
|
||||
slot["tokens"] = token_usage_dict["total_tokens"]
|
||||
slot["prompt_tokens"] = token_usage_dict["prompt_tokens"]
|
||||
slot["completion_tokens"] = token_usage_dict["completion_tokens"]
|
||||
slot["completion"] = result_text or ""
|
||||
slot["initial_completion"] = result_text or ""
|
||||
|
||||
if not result_text:
|
||||
return None, usage
|
||||
|
||||
success, parsed_data, error_msg = parse_json_object_response(
|
||||
result_text, "汇总质量分析"
|
||||
)
|
||||
|
||||
if success and parsed_data:
|
||||
valid, normalized, validation_error = self._validate_review_payload(
|
||||
parsed_data
|
||||
)
|
||||
if valid and normalized:
|
||||
review = self._build_review_from_dict(normalized)
|
||||
logger.info(
|
||||
f"聊天质量汇总分析成功,解析到 {len(review.dimensions)} 个汇总维度"
|
||||
)
|
||||
return review, usage
|
||||
error_msg = validation_error or error_msg
|
||||
|
||||
repaired_data = await self._retry_parse_quality_object(
|
||||
original_prompt=prompt,
|
||||
previous_output=result_text,
|
||||
parse_error=error_msg,
|
||||
umo=umo,
|
||||
system_prompt=system_prompt,
|
||||
base_temperature=base_temperature,
|
||||
)
|
||||
if repaired_data:
|
||||
review = self._build_review_from_dict(repaired_data)
|
||||
logger.info(
|
||||
f"聊天质量汇总 schema 修复重试成功,解析到 {len(review.dimensions)} 个汇总维度"
|
||||
)
|
||||
return review, usage
|
||||
|
||||
# 降级:如果汇总失败,返回最新的一个
|
||||
logger.warning(f"聊天质量汇总分析失败,降级使用最新批次: {error_msg}")
|
||||
return self._build_review_from_dict(batch_reviews[-1]), usage
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"聊天质量汇总分析异常: {e}", exc_info=True)
|
||||
return self._build_review_from_dict(batch_reviews[-1]), TokenUsage()
|
||||
|
||||
async def analyze_quality(
|
||||
self,
|
||||
messages: list[dict],
|
||||
umo: str | None = None,
|
||||
session_id: str | None = None,
|
||||
persona_id: str | None = None,
|
||||
prompt_override: str | None = None,
|
||||
) -> tuple[QualityReview | None, TokenUsage]:
|
||||
"""
|
||||
分析聊天质量
|
||||
|
||||
流程遵循 BaseAnalyzer 的设计模式:
|
||||
1. 构建 prompt
|
||||
2. 调用 LLM
|
||||
3. 提取 token 使用统计
|
||||
4. JSON 解析(使用 parse_json_object_response)
|
||||
5. 正则降级(使用 extract_quality_with_regex)
|
||||
"""
|
||||
try:
|
||||
provider_id_key = self.get_provider_id_key()
|
||||
resolved_provider_id = None
|
||||
if provider_id_key:
|
||||
resolved_provider_id = await get_provider_id_with_fallback(
|
||||
self.context, self.config_manager, provider_id_key, umo
|
||||
)
|
||||
|
||||
# 1. 获取人格设定
|
||||
system_prompt = await self._build_system_prompt(umo)
|
||||
base_temperature = await self._resolve_provider_temperature(
|
||||
self.get_provider_id_key(), umo, provider_id=resolved_provider_id
|
||||
)
|
||||
|
||||
# 2. 构建 prompt
|
||||
prompt = self.build_prompt(messages)
|
||||
if not prompt:
|
||||
return None, TokenUsage()
|
||||
|
||||
# 应用人设强化注入
|
||||
prompt = self._apply_persona_reinforcement(prompt, system_prompt)
|
||||
|
||||
from ....shared.trace_context import TraceContext
|
||||
|
||||
trace = TraceContext.current()
|
||||
if trace:
|
||||
prompts_map = trace.metadata.setdefault("llm_prompts", {})
|
||||
prompts_map[self.get_data_type()] = {
|
||||
"prompt": prompt,
|
||||
"system_prompt": system_prompt,
|
||||
"provider_id": resolved_provider_id or "default",
|
||||
}
|
||||
|
||||
# 3. 调用 LLM
|
||||
response = await call_provider_with_retry(
|
||||
self.context,
|
||||
self.config_manager,
|
||||
prompt=prompt,
|
||||
umo=umo,
|
||||
provider_id_key=self.get_provider_id_key(),
|
||||
provider_id=resolved_provider_id,
|
||||
system_prompt=system_prompt,
|
||||
response_format=self.get_response_format(),
|
||||
observation_label=self.get_data_type(),
|
||||
)
|
||||
|
||||
if response is None:
|
||||
err_text = "聊天质量分析调用LLM失败: Provider 返回空响应或重试耗尽"
|
||||
logger.error(err_text)
|
||||
raise RuntimeError(err_text)
|
||||
|
||||
# 4. 提取 token 使用统计
|
||||
token_usage_dict = extract_token_usage(response)
|
||||
usage = TokenUsage(
|
||||
prompt_tokens=token_usage_dict["prompt_tokens"],
|
||||
completion_tokens=token_usage_dict["completion_tokens"],
|
||||
total_tokens=token_usage_dict["total_tokens"],
|
||||
)
|
||||
|
||||
# 5. 提取响应文本
|
||||
result_text = extract_response_text(response)
|
||||
if not result_text:
|
||||
err_text = "聊天质量分析失败: LLM 未返回任何文本内容"
|
||||
logger.error(err_text)
|
||||
raise RuntimeError(err_text)
|
||||
|
||||
if trace:
|
||||
slot = trace.metadata.setdefault("llm_prompts", {}).setdefault(
|
||||
self.get_data_type(), {}
|
||||
)
|
||||
slot["prompt"] = prompt
|
||||
slot["system_prompt"] = system_prompt
|
||||
slot["provider_id"] = resolved_provider_id or "default"
|
||||
slot["tokens"] = token_usage_dict["total_tokens"]
|
||||
slot["prompt_tokens"] = token_usage_dict["prompt_tokens"]
|
||||
slot["completion_tokens"] = token_usage_dict["completion_tokens"]
|
||||
slot["completion"] = result_text
|
||||
|
||||
# 6. JSON 解析(使用 parse_json_object_response)
|
||||
success, parsed_data, error_msg = parse_json_object_response(
|
||||
result_text, self.get_data_type()
|
||||
)
|
||||
|
||||
if success and parsed_data:
|
||||
valid, normalized, validation_error = self._validate_review_payload(
|
||||
parsed_data
|
||||
)
|
||||
if valid and normalized:
|
||||
review = self._build_review_from_dict(normalized)
|
||||
logger.debug(
|
||||
f"聊天质量分析成功,解析到 {len(review.dimensions)} 个维度"
|
||||
)
|
||||
return review, usage
|
||||
error_msg = validation_error or error_msg
|
||||
|
||||
regex_data = extract_quality_with_regex(result_text)
|
||||
if regex_data:
|
||||
valid, normalized, validation_error = self._validate_review_payload(
|
||||
regex_data
|
||||
)
|
||||
if valid and normalized:
|
||||
review = self._build_review_from_dict(normalized)
|
||||
logger.debug(
|
||||
f"聊天质量首轮结构化失败后,正则提取成功,获得 {len(review.dimensions)} 个维度"
|
||||
)
|
||||
return review, usage
|
||||
error_msg = validation_error or error_msg
|
||||
|
||||
repaired_data = await self._retry_parse_quality_object(
|
||||
original_prompt=prompt,
|
||||
previous_output=result_text,
|
||||
parse_error=error_msg,
|
||||
umo=umo,
|
||||
system_prompt=system_prompt,
|
||||
base_temperature=base_temperature,
|
||||
)
|
||||
if repaired_data:
|
||||
review = self._build_review_from_dict(repaired_data)
|
||||
logger.debug(
|
||||
f"聊天质量 schema 修复重试成功,解析到 {len(review.dimensions)} 个维度"
|
||||
)
|
||||
return review, usage
|
||||
|
||||
# 7. 全部失败
|
||||
err_text = f"聊天质量分析失败: JSON解析和正则提取均未成功 ({error_msg or '未产出有效内容'})"
|
||||
logger.error(err_text)
|
||||
raise RuntimeError(err_text)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"聊天质量分析失败: {e}", exc_info=True)
|
||||
raise
|
||||
|
||||
# Override analyze to bridge the base class interface
|
||||
async def analyze(
|
||||
self,
|
||||
data: list[dict],
|
||||
umo: str | None = None,
|
||||
session_id: str | None = None,
|
||||
persona_id: str | None = None,
|
||||
prompt_override: str | None = None,
|
||||
) -> tuple[list[QualityReview], TokenUsage]:
|
||||
review, usage = await self.analyze_quality(
|
||||
data,
|
||||
umo,
|
||||
session_id,
|
||||
persona_id=persona_id,
|
||||
prompt_override=prompt_override,
|
||||
)
|
||||
return [review] if review else [], usage
|
||||
@@ -0,0 +1,168 @@
|
||||
import re
|
||||
|
||||
from ....domain.models.data_models import TokenUsage
|
||||
from ....utils.logger import logger
|
||||
from ..utils.structured_output_schema import JSONObject
|
||||
from .base_analyzer import BaseAnalyzer
|
||||
|
||||
|
||||
class ComicStoryboardAnalyzer(BaseAnalyzer[dict, list[dict]]):
|
||||
"""
|
||||
分镜及绘画提示词分析器
|
||||
直接从聊天记录中提取金句并生成绘画提示词(含文字渲染要求)
|
||||
"""
|
||||
|
||||
def get_provider_id_key(self) -> str:
|
||||
"""获取画图提示词专用 Provider ID 配置键名"""
|
||||
return "drawing_prompt_provider_id"
|
||||
|
||||
def get_data_type(self) -> str:
|
||||
return "comic_storyboards"
|
||||
|
||||
def get_max_count(self) -> int:
|
||||
return self.config_manager.get_max_topics()
|
||||
|
||||
def build_prompt(self, data: list[dict], prompt_template: str | None = None) -> str:
|
||||
prompt_template = (
|
||||
prompt_template or self.config_manager.get_comic_storyboard_prompt()
|
||||
)
|
||||
if not prompt_template:
|
||||
# 默认的 Prompt
|
||||
prompt_template = (
|
||||
"你是一个资深的漫画分镜师与 AI 绘画提示词专家。\n"
|
||||
"请根据以下给出的【群聊每日核心话题列表】,站在你【当前的人格角色设定】(详见系统注入的身份)的视角与语气,将其改编并设计为一个精彩的多格连环漫画(Comic Strip)全景视觉画图提示词 (Prompt)。\n\n"
|
||||
"【核心视觉、台词与双层排版规则】:\n"
|
||||
"1. 【全话题必须覆盖(共 ${topic_count} 个分格)】:给出的待创作核心话题列表中共有 ${topic_count} 个话题,你必须为每一个话题分别设计一个分格 Panel(即 Panel 1 到 Panel ${topic_count}),绝对不许随意裁减、挑选或遗漏任何一个话题!\n"
|
||||
"2. 【人设口语化台词改编】:绝对严禁将报告分析总结原文直接放进对话框!必须以【你当前的人格语气/性格/口吻】(例如傲娇、萌系或专属说话风格),将每个话题的事件提炼改编为一句角色在漫画中的生动台词或吐槽,【每条台词控制在 15 个汉字以内】(例如:“呜呜!家里云又断网了啦!”、“萝卜子才没有降智!那都是Gemini的错!”)。\n"
|
||||
"3. 【精美双层文字排版(气泡 + 可爱旁白字幕条)】:\n"
|
||||
" - 【角色的气泡】:在描述英文 Prompt 时,指定样式为“adorable kawaii anime speech bubble, soft rounded cloud-like shape, cute pointer tail pointing to the speaker”。\n"
|
||||
" - 【分格底部的事件旁白条】:将每个话题概括为生动精炼的短标题(控制在 30 字以内,严禁带有“【事件】”字样),在 Prompt 中指定样式为“cute pastel-colored kawaii caption strip at the bottom of the panel with soft rounded corners”,严禁死板白框!\n"
|
||||
'4. 【中文文本显式渲染】:画面整体构图与场景描述使用英文,但气泡与底部旁白条内渲染的中文必须显式指定(指令格式:containing a speech bubble with exact Chinese text "人设吐槽台词" 以及 and a cute caption strip at bottom with exact Chinese text "精炼话题短标题"),绝对禁止将中文翻译成英文!\n'
|
||||
"5. 【话题内容直传与文字渲染约束】:在生成传给生图 LLM 的英文提示词 (scene) 时,对于每一个分格,除了描述具体的视觉画面外,你必须将该话题的【完整详情(翻译为英文)】作为 Background Context 附加在该分格的提示词中,帮助生图模型理解剧情。但同时,必须极其强烈地警告生图模型:“绝对禁止将长篇上下文写在画面上,仅允许渲染短标题字幕条和气泡台词!”(示例:Background Context: [Details]. STRICT RULE: DO NOT render the background context text! ONLY render the exact Chinese text in the bubble and caption strip!)。\n"
|
||||
"6. 【核心角色强制全覆盖】:在提示词中必须明确要求并描述,每一个分格 (Panel) 都必须无一例外地出现你当前的人格设定(即参考图中的核心角色,例如 1girl, [特定外貌特征] 等),保持整篇连环画的主角绝对连贯!\n\n"
|
||||
"【待创作的群聊核心话题列表】:\n${chat_content}\n\n"
|
||||
'请输出包含 "scene" 字段的 JSON 对象。\n'
|
||||
)
|
||||
|
||||
valid_topics = [m for m in data if m.get("topic", "")]
|
||||
topic_count = len(valid_topics) if valid_topics else self.get_max_count()
|
||||
chat_content = "\n".join(
|
||||
[
|
||||
f"{i + 1}. 话题: {m.get('topic', '')}\n 详情: {m.get('detail', '')}"
|
||||
for i, m in enumerate(valid_topics)
|
||||
]
|
||||
)
|
||||
|
||||
try:
|
||||
from string import Template
|
||||
|
||||
if "${" in prompt_template or "$" in prompt_template:
|
||||
return Template(prompt_template).safe_substitute(
|
||||
chat_content=chat_content,
|
||||
topic_count=topic_count,
|
||||
max_count=topic_count,
|
||||
)
|
||||
else:
|
||||
return prompt_template.format(
|
||||
chat_content=chat_content,
|
||||
topic_count=topic_count,
|
||||
max_count=topic_count,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(f"漫画分镜提示词格式化失败,使用默认格式: {e}")
|
||||
return f"请从以下群聊话题中提取并生成包含 scene 的 JSON:\n{chat_content}"
|
||||
|
||||
def build_prompt_with_override(
|
||||
self, data: list[dict], prompt_override: str | None
|
||||
) -> str:
|
||||
"""使用角色专属模板构建漫画分镜提示词。"""
|
||||
return self.build_prompt(data, prompt_override)
|
||||
|
||||
def extract_with_regex(self, result_text: str, max_count: int) -> list[dict]:
|
||||
del max_count
|
||||
storyboards = []
|
||||
scene_match = re.search(r'"scene"\s*:\s*"((?:[^"\\]|\\.)*)"', result_text)
|
||||
if scene_match:
|
||||
storyboards.append({"scene": scene_match.group(1).replace('\\"', '"')})
|
||||
return storyboards
|
||||
|
||||
def parse_structured_response(
|
||||
self, result_text: str
|
||||
) -> tuple[bool, list[dict] | None, str | None]:
|
||||
from ..utils.json_utils import parse_json_object_response
|
||||
|
||||
success, data, error = parse_json_object_response(
|
||||
result_text, self.get_data_type()
|
||||
)
|
||||
if success and isinstance(data, dict):
|
||||
if "storyboards" in data and isinstance(data["storyboards"], list):
|
||||
scenes = [
|
||||
item["scene"].strip()
|
||||
for item in data["storyboards"]
|
||||
if isinstance(item, dict)
|
||||
and isinstance(item.get("scene"), str)
|
||||
and item["scene"].strip()
|
||||
]
|
||||
if scenes:
|
||||
return True, [{"scene": "\n\n".join(scenes)}], None
|
||||
return False, None, "'storyboards'中没有有效的非空'scene'字段"
|
||||
elif "scene" in data:
|
||||
scene = data["scene"]
|
||||
if isinstance(scene, str) and scene.strip():
|
||||
return True, [{"scene": scene.strip()}], None
|
||||
return False, None, "'scene'字段必须是非空字符串"
|
||||
else:
|
||||
return False, None, "无法在JSON对象中找到'scene'或'storyboards'字段"
|
||||
return False, None, error
|
||||
|
||||
def create_data_objects(self, data_list: list[dict]) -> list[dict]:
|
||||
# 我们直接返回 dict,因为不需要特别的类型验证
|
||||
return data_list
|
||||
|
||||
def get_response_schema(self) -> JSONObject:
|
||||
return {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"scene": {
|
||||
"type": "string",
|
||||
"description": "One complete panoramic comic image-generation prompt covering every topic and panel",
|
||||
}
|
||||
},
|
||||
"required": ["scene"],
|
||||
"additionalProperties": False,
|
||||
}
|
||||
|
||||
async def analyze_storyboards(
|
||||
self,
|
||||
topics: list[dict],
|
||||
umo: str | None = None,
|
||||
session_id: str | None = None,
|
||||
persona_id: str | None = None,
|
||||
prompt_template: str | None = None,
|
||||
) -> tuple[list[dict], TokenUsage]:
|
||||
"""执行分析,返回 storyboards 和 token 消耗。
|
||||
|
||||
Args:
|
||||
topics: 已提取的有效群聊话题。
|
||||
umo: 群聊统一消息来源标识。
|
||||
session_id: 调试会话标识。
|
||||
persona_id: 漫画分镜专用人格 ID。
|
||||
prompt_template: 角色专属的漫画分镜提示词模板。
|
||||
|
||||
Returns:
|
||||
分镜列表和 Token 使用统计。
|
||||
"""
|
||||
storyboards, usage = await self.analyze(
|
||||
topics, umo, session_id, persona_id, prompt_template
|
||||
)
|
||||
|
||||
if storyboards:
|
||||
if isinstance(storyboards, list):
|
||||
if len(storyboards) > 0 and isinstance(storyboards[0], dict):
|
||||
if "storyboards" in storyboards[0]:
|
||||
storyboards = storyboards[0]["storyboards"]
|
||||
# else: storyboards[0] 已含 "scene",直接使用
|
||||
|
||||
if isinstance(storyboards, list):
|
||||
return [item for item in storyboards if isinstance(item, dict)], usage
|
||||
return [], usage
|
||||
@@ -0,0 +1,222 @@
|
||||
"""
|
||||
金句分析模块
|
||||
专门处理群聊金句提取和分析
|
||||
"""
|
||||
|
||||
from datetime import datetime
|
||||
|
||||
from ....domain.models.data_models import GoldenQuote, TokenUsage
|
||||
from ....utils.logger import logger
|
||||
from ...utils.template_utils import render_template
|
||||
from ..utils import InfoUtils
|
||||
from ..utils.json_utils import extract_golden_quotes_with_regex
|
||||
from ..utils.response_validation import validate_golden_quote_items
|
||||
from ..utils.structured_output_schema import JSONObject, build_golden_quotes_schema
|
||||
from .base_analyzer import BaseAnalyzer
|
||||
|
||||
|
||||
class GoldenQuoteAnalyzer(BaseAnalyzer[GoldenQuote, list[dict]]):
|
||||
"""
|
||||
金句分析器
|
||||
专门处理群聊金句的提取和分析
|
||||
"""
|
||||
|
||||
def get_provider_id_key(self) -> str:
|
||||
"""获取 Provider ID 配置键名"""
|
||||
return "golden_quote_provider_id"
|
||||
|
||||
def get_data_type(self) -> str:
|
||||
"""获取数据类型标识"""
|
||||
return "金句"
|
||||
|
||||
def get_max_count(self) -> int:
|
||||
"""获取最大金句数量,增量模式下使用覆盖值"""
|
||||
if self._incremental_max_count is not None:
|
||||
return self._incremental_max_count
|
||||
return self.config_manager.get_max_golden_quotes()
|
||||
|
||||
def get_response_schema_name(self) -> str:
|
||||
return "daily_golden_quotes"
|
||||
|
||||
def get_response_schema(self) -> JSONObject:
|
||||
return build_golden_quotes_schema(self.get_max_count())
|
||||
|
||||
def build_prompt(self, data: list[dict]) -> str:
|
||||
"""
|
||||
构建金句分析提示词
|
||||
|
||||
Args:
|
||||
messages: 群聊的文本消息列表
|
||||
|
||||
Returns:
|
||||
提示词字符串
|
||||
"""
|
||||
if not data:
|
||||
return ""
|
||||
|
||||
# 构建消息文本 (用 [user_id] 替代 nickname 以确保回填 100% 准确,避免 Emoji 等干扰)
|
||||
messages_text = "\n".join(
|
||||
[f"[{msg['time']}] [{msg['user_id']}]: {msg['content']}" for msg in data]
|
||||
)
|
||||
|
||||
max_golden_quotes = self.get_max_count()
|
||||
|
||||
# 从配置读取 prompt 模板(默认使用 "default" 风格)
|
||||
prompt_template = self.config_manager.get_golden_quote_analysis_prompt()
|
||||
|
||||
if prompt_template:
|
||||
try:
|
||||
prompt = render_template(
|
||||
prompt_template,
|
||||
max_golden_quotes=max_golden_quotes,
|
||||
messages_text=messages_text,
|
||||
)
|
||||
logger.info("使用配置中的金句分析提示词")
|
||||
return prompt
|
||||
except Exception as e:
|
||||
logger.warning(f"应用金句分析提示词失败: {e}")
|
||||
|
||||
logger.warning("未找到有效的金句分析提示词配置,请检查配置文件")
|
||||
return ""
|
||||
|
||||
def extract_with_regex(self, result_text: str, max_count: int) -> list[dict]:
|
||||
"""
|
||||
使用正则表达式提取金句信息
|
||||
|
||||
Args:
|
||||
result_text: LLM响应文本
|
||||
max_count: 最大提取数量
|
||||
|
||||
Returns:
|
||||
金句数据列表
|
||||
"""
|
||||
return extract_golden_quotes_with_regex(result_text, max_count)
|
||||
|
||||
def create_data_objects(self, data_list: list[dict]) -> list[GoldenQuote]:
|
||||
"""
|
||||
创建金句对象列表
|
||||
|
||||
Args:
|
||||
quotes_data: 原始金句数据列表
|
||||
|
||||
Returns:
|
||||
GoldenQuote对象列表
|
||||
"""
|
||||
try:
|
||||
quotes = []
|
||||
max_quotes = self.get_max_count()
|
||||
|
||||
for quote_data in data_list[:max_quotes]:
|
||||
# 确保数据格式正确
|
||||
content = quote_data.get("content", "").strip()
|
||||
sender = quote_data.get("sender", "").strip()
|
||||
reason = quote_data.get("reason", "").strip()
|
||||
|
||||
# 验证必要字段
|
||||
if not content or not sender or not reason:
|
||||
logger.warning(f"金句数据格式不完整,跳过: {quote_data}")
|
||||
continue
|
||||
|
||||
quotes.append(
|
||||
GoldenQuote(content=content, sender=sender, reason=reason)
|
||||
)
|
||||
|
||||
return quotes
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"创建金句对象失败: {e}")
|
||||
return []
|
||||
|
||||
def validate_parsed_data(
|
||||
self, data_list: list[dict]
|
||||
) -> tuple[bool, list[dict] | None, str | None]:
|
||||
return validate_golden_quote_items(data_list)
|
||||
|
||||
async def analyze_golden_quotes(
|
||||
self,
|
||||
messages: list[dict],
|
||||
umo: str | None = None,
|
||||
session_id: str | None = None,
|
||||
) -> tuple[list[GoldenQuote], TokenUsage]:
|
||||
"""
|
||||
分析群聊金句
|
||||
|
||||
Args:
|
||||
messages: 群聊消息列表
|
||||
umo: 模型唯一标识符
|
||||
session_id: 会话ID (用于调试模式)
|
||||
|
||||
Returns:
|
||||
(金句列表, Token使用统计)
|
||||
"""
|
||||
try:
|
||||
# 提取圣经的文本消息
|
||||
interesting_messages = self.extract_interesting_messages(messages)
|
||||
|
||||
if not interesting_messages:
|
||||
logger.info("没有符合条件的圣经消息,返回空结果")
|
||||
return [], TokenUsage()
|
||||
|
||||
logger.info(f"开始从 {len(interesting_messages)} 条圣经消息中提取金句")
|
||||
quotes, usage = await self.analyze(interesting_messages, umo, session_id)
|
||||
|
||||
# 建立 ID 到昵称的映射表用于恢复显示
|
||||
id_to_nickname = {}
|
||||
for msg in interesting_messages:
|
||||
uid = str(msg.get("user_id", ""))
|
||||
if uid:
|
||||
id_to_nickname[uid] = msg.get("sender", "")
|
||||
|
||||
# 回填 User ID 并恢复发送者昵称
|
||||
for quote in quotes:
|
||||
# 此时 quote.sender 包含的是 Prompt 中的 [user_id]
|
||||
# 有些 LLM 可能会带上中括号,尝试清理
|
||||
potential_id = quote.sender.strip().strip("[]")
|
||||
|
||||
if potential_id in id_to_nickname:
|
||||
quote.user_id = potential_id
|
||||
quote.sender = id_to_nickname[potential_id]
|
||||
else:
|
||||
logger.warning(
|
||||
f"[金句分析] 无法匹配 User ID: {potential_id},金句将无法显示真实头像。"
|
||||
)
|
||||
|
||||
return quotes, usage
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"金句分析失败: {e}")
|
||||
raise
|
||||
|
||||
def extract_interesting_messages(self, messages: list[dict]) -> list[dict]:
|
||||
"""
|
||||
根据清理后的消息提取可能有意义的消息片段用于金句分析。
|
||||
|
||||
Args:
|
||||
messages: 已由 MessageCleaner 处理过的 legacy 消息列表
|
||||
|
||||
Returns:
|
||||
提取的文本消息列表
|
||||
"""
|
||||
interesting_messages = []
|
||||
|
||||
for msg in messages:
|
||||
# 获取发送者显示名
|
||||
sender = msg.get("sender", {})
|
||||
nickname = InfoUtils.get_user_nickname(self.config_manager, sender)
|
||||
msg_time = datetime.fromtimestamp(msg.get("time", 0)).strftime("%H:%M")
|
||||
|
||||
for content in msg.get("message", []):
|
||||
if content.get("type") == "text":
|
||||
text = content.get("data", {}).get("text", "").strip()
|
||||
# 过滤掉过短或过长的噪音(已经在 cleaner 处理过一遍基本垃圾)
|
||||
if 2 <= len(text) <= 500:
|
||||
interesting_messages.append(
|
||||
{
|
||||
"sender": nickname,
|
||||
"time": msg_time,
|
||||
"content": text,
|
||||
"user_id": str(sender.get("user_id", "")),
|
||||
}
|
||||
)
|
||||
|
||||
return interesting_messages
|
||||
@@ -0,0 +1,395 @@
|
||||
"""
|
||||
话题分析模块
|
||||
专门处理群聊话题分析
|
||||
"""
|
||||
|
||||
import re
|
||||
from datetime import datetime
|
||||
|
||||
from ....domain.models.data_models import SummaryTopic, TokenUsage
|
||||
from ....utils.logger import logger
|
||||
from ...utils.template_utils import render_template
|
||||
from ..utils import InfoUtils
|
||||
from ..utils.json_utils import extract_topics_with_regex
|
||||
from ..utils.response_validation import validate_topic_items
|
||||
from ..utils.structured_output_schema import JSONObject, build_topics_schema
|
||||
from .base_analyzer import BaseAnalyzer
|
||||
|
||||
|
||||
class TopicAnalyzer(BaseAnalyzer[SummaryTopic, list[dict]]):
|
||||
"""
|
||||
话题分析器
|
||||
专门处理群聊话题的提取和分析
|
||||
"""
|
||||
|
||||
def get_provider_id_key(self) -> str:
|
||||
"""获取 Provider ID 配置键名"""
|
||||
return "topic_provider_id"
|
||||
|
||||
def get_data_type(self) -> str:
|
||||
"""获取数据类型标识"""
|
||||
return "话题"
|
||||
|
||||
def get_max_count(self) -> int:
|
||||
"""获取最大话题数量,增量模式下使用覆盖值"""
|
||||
if self._incremental_max_count is not None:
|
||||
return self._incremental_max_count
|
||||
return self.config_manager.get_max_topics()
|
||||
|
||||
def get_response_schema_name(self) -> str:
|
||||
return "daily_topics"
|
||||
|
||||
def get_response_schema(self) -> JSONObject:
|
||||
return build_topics_schema(self.get_max_count())
|
||||
|
||||
def build_prompt(self, data: list[dict]) -> str:
|
||||
"""
|
||||
构建话题分析提示词
|
||||
|
||||
Args:
|
||||
messages: 群聊消息列表
|
||||
|
||||
Returns:
|
||||
提示词字符串
|
||||
"""
|
||||
# 验证输入数据格式
|
||||
if not isinstance(data, list):
|
||||
logger.error(f"build_prompt 期望列表,但收到: {type(data)}")
|
||||
return ""
|
||||
|
||||
# 检查消息列表是否为空
|
||||
if not data:
|
||||
logger.warning("build_prompt 收到空消息列表")
|
||||
return ""
|
||||
|
||||
# 提取文本消息
|
||||
text_messages = []
|
||||
for i, msg in enumerate(data):
|
||||
# 确保msg是字典类型,避免'str' object has no attribute 'get'错误
|
||||
if not isinstance(msg, dict):
|
||||
continue
|
||||
|
||||
try:
|
||||
sender = msg.get("sender", {})
|
||||
# 确保sender是字典类型,避免'str' object has no attribute 'get'错误
|
||||
if not isinstance(sender, dict):
|
||||
continue
|
||||
|
||||
# 获取发送者ID并过滤机器人消息
|
||||
user_id = str(sender.get("user_id", ""))
|
||||
bot_self_ids = self.config_manager.get_bot_self_ids()
|
||||
|
||||
# 跳过机器人自己的消息
|
||||
if bot_self_ids and user_id in [str(uid) for uid in bot_self_ids]:
|
||||
continue
|
||||
|
||||
nickname = InfoUtils.get_user_nickname(self.config_manager, sender)
|
||||
msg_time = datetime.fromtimestamp(msg.get("time", 0)).strftime("%H:%M")
|
||||
|
||||
message_list = msg.get("message", [])
|
||||
|
||||
# 提取文本内容,可能分布在多个 content 中
|
||||
text_parts = []
|
||||
for j, content in enumerate(message_list):
|
||||
if not isinstance(content, dict):
|
||||
continue
|
||||
|
||||
content_type = content.get("type", "")
|
||||
|
||||
if content_type == "text":
|
||||
text = content.get("data", {}).get("text", "").strip()
|
||||
if text:
|
||||
text_parts.append(text)
|
||||
elif content_type == "at":
|
||||
# 处理 @ 消息,转换为文本
|
||||
at_data = content.get("data", {})
|
||||
# 兼容不同平台的 ID 字段
|
||||
at_id = at_data.get("id") or at_data.get("user_id")
|
||||
if at_id:
|
||||
at_text = f"@{at_id}"
|
||||
text_parts.append(at_text)
|
||||
elif content_type == "reply":
|
||||
# 处理回复消息,添加标记
|
||||
reply_id = content.get("data", {}).get("id", "")
|
||||
if reply_id:
|
||||
reply_text = f"[回复:{reply_id}]"
|
||||
text_parts.append(reply_text)
|
||||
|
||||
# 合并所有文本部分
|
||||
combined_text = "".join(text_parts).strip()
|
||||
|
||||
if (
|
||||
combined_text
|
||||
and len(combined_text) > 2
|
||||
and not combined_text.startswith("/")
|
||||
):
|
||||
# 清理消息内容
|
||||
cleaned_text = combined_text.replace("“", '"').replace("”", '"')
|
||||
cleaned_text = cleaned_text.replace("‘", "'").replace("’", "'")
|
||||
cleaned_text = cleaned_text.replace("\n", " ").replace("\r", " ")
|
||||
cleaned_text = cleaned_text.replace("\t", " ")
|
||||
cleaned_text = re.sub(r"[\x00-\x1f\x7f-\x9f]", "", cleaned_text)
|
||||
|
||||
text_messages.append(
|
||||
{
|
||||
"sender": nickname,
|
||||
"time": msg_time,
|
||||
"content": cleaned_text,
|
||||
"user_id": str(user_id),
|
||||
}
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"build_prompt 处理第 {i + 1} 条消息时出错: {e}", exc_info=True
|
||||
)
|
||||
continue
|
||||
|
||||
if not text_messages:
|
||||
logger.warning("build_prompt 没有提取到有效的文本消息,返回空prompt")
|
||||
return ""
|
||||
|
||||
# 构建消息文本
|
||||
# 使用用户提供的 ID-Only 格式: [HH:MM] [用户ID]: 消息内容
|
||||
messages_text = "\n".join(
|
||||
[
|
||||
f"[{msg['time']}] [{msg['user_id']}]: {msg['content']}"
|
||||
for msg in text_messages
|
||||
]
|
||||
)
|
||||
|
||||
max_topics = self.get_max_count()
|
||||
|
||||
# 从配置读取 prompt 模板(默认使用 "default" 风格)
|
||||
prompt_template = self.config_manager.get_topic_analysis_prompt()
|
||||
|
||||
if prompt_template:
|
||||
try:
|
||||
prompt = render_template(
|
||||
prompt_template,
|
||||
max_topics=max_topics,
|
||||
messages_text=messages_text,
|
||||
)
|
||||
logger.info("使用配置中的话题分析提示词")
|
||||
return prompt
|
||||
except Exception as e:
|
||||
logger.warning(f"应用话题分析提示词失败: {e}")
|
||||
|
||||
logger.warning("未找到有效的话题分析提示词配置,请检查配置文件")
|
||||
return ""
|
||||
|
||||
def extract_with_regex(self, result_text: str, max_count: int) -> list[dict]:
|
||||
"""
|
||||
使用正则表达式提取话题信息
|
||||
|
||||
Args:
|
||||
result_text: LLM响应文本
|
||||
max_count: 最大话题数量
|
||||
|
||||
Returns:
|
||||
话题数据列表
|
||||
"""
|
||||
return extract_topics_with_regex(result_text, max_count)
|
||||
|
||||
def create_data_objects(self, data_list: list[dict]) -> list[SummaryTopic]:
|
||||
"""
|
||||
创建话题对象列表
|
||||
|
||||
Args:
|
||||
topics_data: 原始话题数据列表
|
||||
|
||||
Returns:
|
||||
SummaryTopic对象列表
|
||||
"""
|
||||
logger.debug(
|
||||
f"create_data_objects 开始处理,输入数据数量: {len(data_list) if data_list else 0}"
|
||||
)
|
||||
logger.debug(f"输入数据类型: {type(data_list)}")
|
||||
|
||||
try:
|
||||
topics = []
|
||||
max_topics = self.get_max_count()
|
||||
|
||||
logger.debug(f"处理前 {max_topics} 条话题数据")
|
||||
|
||||
for i, topic_data in enumerate(data_list[:max_topics]):
|
||||
logger.debug(f"处理第 {i + 1} 条话题数据,类型: {type(topic_data)}")
|
||||
|
||||
# 确保topic_data是字典类型,避免'str' object has no attribute 'get'错误
|
||||
if not isinstance(topic_data, dict):
|
||||
logger.warning(
|
||||
f"跳过非字典类型的话题数据: {type(topic_data)} - {topic_data}"
|
||||
)
|
||||
continue
|
||||
|
||||
try:
|
||||
# 确保数据格式正确
|
||||
topic_name = topic_data.get("topic", "").strip()
|
||||
contributors = topic_data.get("contributors", [])
|
||||
detail = topic_data.get("detail", "").strip()
|
||||
|
||||
logger.debug(
|
||||
f"话题数据 - 名称: {topic_name}, 参与者: {contributors}, 详情: {detail[:50]}..."
|
||||
)
|
||||
|
||||
# 验证必要字段
|
||||
if not topic_name or not detail:
|
||||
logger.warning(f"话题数据格式不完整,跳过: {topic_data}")
|
||||
continue
|
||||
|
||||
# 确保参与者列表有效
|
||||
if not contributors or not isinstance(contributors, list):
|
||||
contributors = ["群友"]
|
||||
else:
|
||||
# 清理参与者名称
|
||||
contributors = [
|
||||
str(c).strip() for c in contributors if c and str(c).strip()
|
||||
] or ["群友"]
|
||||
|
||||
topics.append(
|
||||
SummaryTopic(
|
||||
topic=topic_name,
|
||||
contributors=contributors[:5], # 最多5个参与者
|
||||
detail=detail,
|
||||
)
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f"处理第 {i + 1} 条话题数据时出错: {e}", exc_info=True)
|
||||
continue
|
||||
|
||||
logger.debug(f"create_data_objects 完成,创建了 {len(topics)} 个话题对象")
|
||||
return topics
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"创建话题对象失败: {e}", exc_info=True)
|
||||
return []
|
||||
|
||||
def validate_parsed_data(
|
||||
self, data_list: list[dict]
|
||||
) -> tuple[bool, list[dict] | None, str | None]:
|
||||
return validate_topic_items(data_list)
|
||||
|
||||
def extract_text_messages(self, messages: list[dict]) -> list[dict]:
|
||||
"""
|
||||
从已清理的消息中提取文本消息用于话题分析。
|
||||
|
||||
Args:
|
||||
messages: 已由 MessageCleaner 处理过的 legacy 消息列表
|
||||
|
||||
Returns:
|
||||
提取的文本消息列表
|
||||
"""
|
||||
text_messages = []
|
||||
|
||||
for msg in messages:
|
||||
# 获取发送者显示名
|
||||
sender = msg.get("sender", {})
|
||||
nickname = InfoUtils.get_user_nickname(self.config_manager, sender)
|
||||
msg_time = datetime.fromtimestamp(msg.get("time", 0)).strftime("%H:%M")
|
||||
|
||||
for content in msg.get("message", []):
|
||||
if content.get("type") == "text":
|
||||
text = content.get("data", {}).get("text", "").strip()
|
||||
# 已经在 MessageCleaner 中处理过基本的垃圾内容
|
||||
if text:
|
||||
# 简单的额外清理
|
||||
cleaned_text = text.replace("\n", " ").replace("\r", " ")
|
||||
cleaned_text = re.sub(r"[\x00-\x1f\x7f-\x9f]", "", cleaned_text)
|
||||
|
||||
text_messages.append(
|
||||
{
|
||||
"sender": nickname,
|
||||
"time": msg_time,
|
||||
"content": cleaned_text.strip(),
|
||||
"user_id": str(sender.get("user_id", "")),
|
||||
}
|
||||
)
|
||||
return text_messages
|
||||
|
||||
async def analyze_topics(
|
||||
self,
|
||||
messages: list[dict],
|
||||
umo: str | None = None,
|
||||
session_id: str | None = None,
|
||||
) -> tuple[list[SummaryTopic], TokenUsage]:
|
||||
"""
|
||||
分析群聊话题
|
||||
|
||||
Args:
|
||||
messages: 群聊消息列表
|
||||
umo: 模型唯一标识符
|
||||
session_id: 会话ID (用于调试模式)
|
||||
|
||||
Returns:
|
||||
(话题列表, Token使用统计)
|
||||
"""
|
||||
try:
|
||||
logger.debug(
|
||||
f"analyze_topics 开始处理,消息数量: {len(messages) if messages else 0}"
|
||||
)
|
||||
logger.debug(f"消息类型: {type(messages)}")
|
||||
if messages:
|
||||
logger.debug(
|
||||
f"第一条消息类型: {type(messages[0]) if messages else '无'}"
|
||||
)
|
||||
logger.debug(f"第一条消息内容: {messages[0] if messages else '无'}")
|
||||
|
||||
# 检查是否有有效的文本消息
|
||||
text_messages = self.extract_text_messages(messages)
|
||||
logger.debug(f"提取到 {len(text_messages)} 条文本消息")
|
||||
|
||||
if not text_messages:
|
||||
logger.info("没有有效的文本消息,返回空结果")
|
||||
return [], TokenUsage()
|
||||
|
||||
logger.info(f"开始分析 {len(text_messages)} 条文本消息中的话题")
|
||||
logger.debug(f"文本消息类型: {type(text_messages)}")
|
||||
if text_messages:
|
||||
logger.debug(f"第一条文本消息类型: {type(text_messages[0])}")
|
||||
logger.debug(f"第一条文本消息内容: {text_messages[0]}")
|
||||
|
||||
# 建立 ID 到昵称的映射表
|
||||
id_to_nickname = {}
|
||||
for msg in text_messages:
|
||||
sender = msg.get("sender")
|
||||
user_id = msg.get("user_id")
|
||||
if sender and user_id:
|
||||
id_to_nickname[user_id] = sender
|
||||
|
||||
# 直接传入原始消息,让 build_prompt 方法处理
|
||||
topics, usage = await self.analyze(messages, umo, session_id)
|
||||
|
||||
# 后处理:contributors 此时包含的是 ID,需要映射回昵称
|
||||
for topic in topics:
|
||||
raw_ids = topic.contributors # LLM 返回的是 ID 列表
|
||||
|
||||
# 填充 contributor_ids。QQ 官方 member_openid 并非纯数字,
|
||||
# 因此仅接受本批次已知用户或已配置机器人 ID,而不是用 isdigit 过滤。
|
||||
bot_ids = {str(uid) for uid in self.config_manager.get_bot_self_ids()}
|
||||
known_ids = set(id_to_nickname) | bot_ids
|
||||
valid_ids = []
|
||||
for raw_uid in raw_ids:
|
||||
uid = str(raw_uid).strip().strip("[]")
|
||||
if uid and uid in known_ids and uid not in valid_ids:
|
||||
valid_ids.append(uid)
|
||||
topic.contributor_ids = valid_ids
|
||||
|
||||
# 映射回昵称用于显示
|
||||
resolved_names = []
|
||||
for uid in valid_ids:
|
||||
# 尝试从当前批次消息映射
|
||||
name = id_to_nickname.get(uid)
|
||||
if not name:
|
||||
# 尝试去全局配置里找 (e.g. 机器人自己)
|
||||
if uid in bot_ids:
|
||||
name = "Bot"
|
||||
else:
|
||||
name = uid # Fallback to ID
|
||||
resolved_names.append(name)
|
||||
|
||||
topic.contributors = resolved_names
|
||||
|
||||
return topics, usage
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"话题分析失败: {e}", exc_info=True)
|
||||
raise
|
||||
@@ -0,0 +1,275 @@
|
||||
"""
|
||||
用户称号分析模块
|
||||
专门处理用户称号和MBTI类型分析
|
||||
"""
|
||||
|
||||
from ....domain.models.data_models import TokenUsage, UserTitle
|
||||
from ....utils.logger import logger
|
||||
from ...utils.template_utils import render_template
|
||||
from ..utils.json_utils import extract_user_titles_with_regex
|
||||
from ..utils.response_validation import validate_user_title_items
|
||||
from ..utils.structured_output_schema import JSONObject, build_user_titles_schema
|
||||
from .base_analyzer import BaseAnalyzer
|
||||
|
||||
|
||||
class UserTitleAnalyzer(BaseAnalyzer[UserTitle, dict]):
|
||||
"""
|
||||
用户称号分析器
|
||||
专门处理用户称号分配和MBTI类型分析
|
||||
"""
|
||||
|
||||
def get_provider_id_key(self) -> str:
|
||||
"""获取 Provider ID 配置键名"""
|
||||
return "user_title_provider_id"
|
||||
|
||||
def get_data_type(self) -> str:
|
||||
"""获取数据类型标识"""
|
||||
return "用户称号"
|
||||
|
||||
def get_max_count(self) -> int:
|
||||
"""获取最大用户称号数量"""
|
||||
return self.config_manager.get_max_user_titles()
|
||||
|
||||
def get_response_schema_name(self) -> str:
|
||||
return "daily_user_titles"
|
||||
|
||||
def get_response_schema(self) -> JSONObject:
|
||||
return build_user_titles_schema(self.get_max_count())
|
||||
|
||||
def build_prompt(self, data: dict) -> str:
|
||||
"""
|
||||
构建用户称号分析提示词
|
||||
|
||||
Args:
|
||||
user_data: 用户数据字典,包含用户统计信息
|
||||
|
||||
Returns:
|
||||
提示词字符串
|
||||
"""
|
||||
user_summaries = data.get("user_summaries", [])
|
||||
|
||||
if not user_summaries:
|
||||
return ""
|
||||
|
||||
# 构建用户数据文本
|
||||
users_text = "\n".join(
|
||||
[
|
||||
f"- {user['name']} (ID:{user['user_id']}): "
|
||||
f"发言{user['message_count']}条, 平均{user['avg_chars']}字, "
|
||||
f"表情比例{user['emoji_ratio']}, 夜间发言比例{user['night_ratio']}, "
|
||||
f"回复比例{user['reply_ratio']}"
|
||||
for user in user_summaries
|
||||
]
|
||||
)
|
||||
|
||||
# 从配置读取 prompt 模板(默认使用 "default" 风格)
|
||||
prompt_template = self.config_manager.get_user_title_analysis_prompt()
|
||||
|
||||
if prompt_template:
|
||||
try:
|
||||
prompt = render_template(prompt_template, users_text=users_text)
|
||||
logger.info("使用配置中的用户称号分析提示词")
|
||||
return prompt
|
||||
except Exception as e:
|
||||
logger.warning(f"应用用户称号分析提示词失败: {e}")
|
||||
|
||||
logger.warning("未找到有效的用户称号分析提示词配置,请检查配置文件")
|
||||
return ""
|
||||
|
||||
def extract_with_regex(self, result_text: str, max_count: int) -> list[dict]:
|
||||
"""
|
||||
使用正则表达式提取用户称号信息
|
||||
|
||||
Args:
|
||||
result_text: LLM响应文本
|
||||
max_count: 最大提取数量
|
||||
|
||||
Returns:
|
||||
用户称号数据列表
|
||||
"""
|
||||
return extract_user_titles_with_regex(result_text, max_count)
|
||||
|
||||
def create_data_objects(self, data_list: list[dict]) -> list[UserTitle]:
|
||||
"""
|
||||
创建用户称号对象列表
|
||||
|
||||
Args:
|
||||
titles_data: 原始用户称号数据列表
|
||||
|
||||
Returns:
|
||||
UserTitle对象列表
|
||||
"""
|
||||
try:
|
||||
titles = []
|
||||
max_titles = self.get_max_count()
|
||||
|
||||
for title_data in data_list[:max_titles]:
|
||||
# 确保数据格式正确
|
||||
name = title_data.get("name", "").strip()
|
||||
user_id = title_data.get("user_id")
|
||||
title = title_data.get("title", "").strip()
|
||||
mbti = title_data.get("mbti", "").strip()
|
||||
reason = title_data.get("reason", "").strip()
|
||||
|
||||
# 验证必要字段
|
||||
if not name or not title or not mbti or not reason:
|
||||
logger.warning(f"用户称号数据格式不完整,跳过: {title_data}")
|
||||
continue
|
||||
|
||||
# 确保 user_id 是字符串
|
||||
if user_id is not None:
|
||||
user_id = str(user_id)
|
||||
else:
|
||||
logger.warning(f"未找到用户ID (user_id),跳过: {title_data}")
|
||||
continue
|
||||
|
||||
titles.append(
|
||||
UserTitle(
|
||||
name=name,
|
||||
user_id=user_id,
|
||||
title=title,
|
||||
mbti=mbti,
|
||||
reason=reason,
|
||||
)
|
||||
)
|
||||
|
||||
return titles
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"创建用户称号对象失败: {e}")
|
||||
return []
|
||||
|
||||
def validate_parsed_data(
|
||||
self, data_list: list[dict]
|
||||
) -> tuple[bool, list[dict] | None, str | None]:
|
||||
return validate_user_title_items(data_list)
|
||||
|
||||
def prepare_user_data(
|
||||
self,
|
||||
messages: list[dict],
|
||||
user_analysis: dict,
|
||||
top_users: list[dict] | None = None,
|
||||
) -> dict:
|
||||
"""
|
||||
准备用户数据
|
||||
|
||||
Args:
|
||||
messages: 群聊消息列表
|
||||
user_analysis: 用户分析统计
|
||||
top_users: 活跃用户列表(从get_top_users获取)
|
||||
|
||||
Returns:
|
||||
准备好的用户数据字典
|
||||
"""
|
||||
try:
|
||||
# 获取机器人 ID 列表用于过滤
|
||||
bot_self_ids = self.config_manager.get_bot_self_ids()
|
||||
|
||||
user_summaries = []
|
||||
|
||||
# 如果提供了top_users列表,只分析这些活跃用户
|
||||
if top_users:
|
||||
logger.info(
|
||||
f"使用get_top_users筛选出的 {len(top_users)} 个活跃用户进行称号分析"
|
||||
)
|
||||
target_user_ids = {str(user["user_id"]) for user in top_users}
|
||||
else:
|
||||
# 兼容旧逻辑:如果没有提供top_users,则使用所有消息数>=5的用户
|
||||
logger.info("未提供活跃用户列表,使用消息数>=5的用户")
|
||||
target_user_ids = {
|
||||
user_id
|
||||
for user_id, stats in user_analysis.items()
|
||||
if stats["message_count"] >= 5
|
||||
}
|
||||
|
||||
for user_id, stats in user_analysis.items():
|
||||
user_id_str = str(user_id)
|
||||
# 过滤机器人由 MessageCleaner 已处理,此处仅作为二级防御
|
||||
if bot_self_ids and user_id_str in [str(uid) for uid in bot_self_ids]:
|
||||
continue
|
||||
|
||||
# 只处理活跃用户 (top_users 或 消息数>=5)
|
||||
if user_id_str not in target_user_ids:
|
||||
continue
|
||||
|
||||
# 分析用户特征 (此处已基于已清理的 stats)
|
||||
# 兼容性处理:优先使用 hours (dict),如果没有则尝试从消息推断或使用空
|
||||
hours_data = stats.get("hours")
|
||||
if hours_data is None:
|
||||
# 尝试兼容旧 schema 或简化版
|
||||
active_hours = stats.get("active_hours", [])
|
||||
hours_data = dict.fromkeys(active_hours, 1)
|
||||
|
||||
# 安全计算夜间发言数
|
||||
night_messages = sum(hours_data.get(h, 0) for h in range(6))
|
||||
|
||||
message_count = stats.get("message_count", 0)
|
||||
if message_count <= 0:
|
||||
continue
|
||||
|
||||
avg_chars = stats.get("char_count", 0) / message_count
|
||||
|
||||
# 称号所需维度
|
||||
user_summaries.append(
|
||||
{
|
||||
"name": stats.get("nickname", stats.get("name", user_id_str)),
|
||||
"user_id": user_id_str,
|
||||
"message_count": message_count,
|
||||
"avg_chars": round(avg_chars, 1),
|
||||
"emoji_ratio": round(
|
||||
stats.get("emoji_count", 0) / message_count, 2
|
||||
),
|
||||
"night_ratio": round(night_messages / message_count, 2),
|
||||
"reply_ratio": round(
|
||||
stats.get("reply_count", 0) / message_count, 2
|
||||
),
|
||||
}
|
||||
)
|
||||
|
||||
if not user_summaries:
|
||||
return {"user_summaries": []}
|
||||
|
||||
# 按消息数量排序
|
||||
user_summaries.sort(key=lambda x: x["message_count"], reverse=True)
|
||||
|
||||
return {"user_summaries": user_summaries}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"准备用户数据失败: {e}")
|
||||
return {"user_summaries": []}
|
||||
|
||||
async def analyze_user_titles(
|
||||
self,
|
||||
messages: list[dict],
|
||||
user_activity: dict,
|
||||
umo: str | None = None,
|
||||
top_users: list[dict] | None = None,
|
||||
session_id: str | None = None,
|
||||
) -> tuple[list[UserTitle], TokenUsage]:
|
||||
"""
|
||||
分析用户称号
|
||||
|
||||
Args:
|
||||
messages: 群聊消息列表
|
||||
user_analysis: 用户分析统计
|
||||
umo: 模型唯一标识符
|
||||
top_users: 活跃用户列表(从get_top_users获取,可选)
|
||||
session_id: 会话ID (用于调试模式)
|
||||
|
||||
Returns:
|
||||
(用户称号列表, Token使用统计)
|
||||
"""
|
||||
try:
|
||||
# 准备用户数据,传入活跃用户列表
|
||||
user_data = self.prepare_user_data(messages, user_activity, top_users)
|
||||
|
||||
if not user_data["user_summaries"]:
|
||||
logger.info("没有符合条件的用户,返回空结果")
|
||||
return [], TokenUsage()
|
||||
|
||||
logger.info(f"开始分析 {len(user_data['user_summaries'])} 个活跃用户的称号")
|
||||
return await self.analyze(user_data, umo, session_id)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"用户称号分析失败: {e}")
|
||||
raise
|
||||
@@ -0,0 +1,802 @@
|
||||
"""
|
||||
LLM分析器模块
|
||||
负责协调各个分析器进行话题分析、用户称号分析和金句分析
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
|
||||
from ...domain.models.data_models import (
|
||||
GoldenQuote,
|
||||
QualityReview,
|
||||
SummaryTopic,
|
||||
TokenUsage,
|
||||
UserTitle,
|
||||
)
|
||||
from ...domain.repositories.analysis_repository import IAnalysisProvider
|
||||
from ...shared.constants import PLUGIN_NAME
|
||||
from ...shared.trace_context import TraceContext
|
||||
from ...utils.logger import logger
|
||||
from .analyzers.chat_quality_analyzer import ChatQualityAnalyzer
|
||||
from .analyzers.comic_analyzer import ComicStoryboardAnalyzer
|
||||
from .analyzers.golden_quote_analyzer import GoldenQuoteAnalyzer
|
||||
from .analyzers.topic_analyzer import TopicAnalyzer
|
||||
from .analyzers.user_title_analyzer import UserTitleAnalyzer
|
||||
from .utils.json_utils import fix_json
|
||||
from .utils.llm_utils import call_provider_with_retry
|
||||
|
||||
|
||||
class LLMAnalyzer(IAnalysisProvider):
|
||||
"""
|
||||
LLM分析器
|
||||
作为统一入口,协调各个专门的分析器进行不同类型的分析
|
||||
保持向后兼容性,提供原有的接口
|
||||
"""
|
||||
|
||||
topic_analyzer: TopicAnalyzer
|
||||
user_title_analyzer: UserTitleAnalyzer
|
||||
golden_quote_analyzer: GoldenQuoteAnalyzer
|
||||
comic_storyboard_analyzer: ComicStoryboardAnalyzer
|
||||
|
||||
def __init__(self, context, config_manager):
|
||||
"""
|
||||
初始化LLM分析器
|
||||
|
||||
Args:
|
||||
context: AstrBot上下文对象
|
||||
config_manager: 配置管理器
|
||||
"""
|
||||
self.context = context
|
||||
self.config_manager = config_manager
|
||||
|
||||
# 初始化各个专门的分析器
|
||||
self.topic_analyzer = TopicAnalyzer(context, config_manager)
|
||||
self.user_title_analyzer = UserTitleAnalyzer(context, config_manager)
|
||||
self.golden_quote_analyzer = GoldenQuoteAnalyzer(context, config_manager)
|
||||
self.chat_quality_analyzer = ChatQualityAnalyzer(context, config_manager)
|
||||
self.comic_storyboard_analyzer = ComicStoryboardAnalyzer(
|
||||
context, config_manager
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _make_session_id(
|
||||
session_id: str | None, umo: str | None = None, prefix: str = ""
|
||||
) -> str:
|
||||
"""Generate a session ID if not already provided."""
|
||||
if session_id:
|
||||
return session_id
|
||||
from datetime import datetime
|
||||
|
||||
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||
if umo:
|
||||
safe_umo = umo.replace(":", "_")
|
||||
return f"{prefix}{timestamp}_{safe_umo}"
|
||||
return f"{prefix}{timestamp}"
|
||||
|
||||
async def analyze_topics(
|
||||
self,
|
||||
messages: list[dict],
|
||||
umo: str | None = None,
|
||||
session_id: str | None = None,
|
||||
) -> tuple[list[SummaryTopic], TokenUsage]:
|
||||
"""
|
||||
使用LLM分析话题
|
||||
保持原有接口,委托给专门的TopicAnalyzer处理
|
||||
|
||||
Args:
|
||||
messages: 群聊消息列表
|
||||
umo: 模型唯一标识符
|
||||
session_id: 会话ID (用于调试模式)
|
||||
|
||||
Returns:
|
||||
(话题列表, Token使用统计)
|
||||
"""
|
||||
try:
|
||||
session_id = self._make_session_id(session_id, umo)
|
||||
|
||||
logger.info(f"开始话题分析, session_id: {session_id}")
|
||||
return await self.topic_analyzer.analyze_topics(messages, umo, session_id)
|
||||
except Exception as e:
|
||||
logger.error(f"话题分析失败: {e}")
|
||||
return [], TokenUsage()
|
||||
|
||||
async def analyze_user_titles(
|
||||
self,
|
||||
messages: list[dict],
|
||||
user_activity: dict,
|
||||
umo: str | None = None,
|
||||
top_users: list[dict] | None = None,
|
||||
session_id: str | None = None,
|
||||
) -> tuple[list[UserTitle], TokenUsage]:
|
||||
"""
|
||||
使用LLM分析用户称号
|
||||
保持原有接口,委托给专门的UserTitleAnalyzer处理
|
||||
|
||||
Args:
|
||||
messages: 群聊消息列表
|
||||
user_activity: 用户分析统计
|
||||
umo: 模型唯一标识符
|
||||
top_users: 活跃用户列表(可选)
|
||||
session_id: 会话ID (用于调试模式)
|
||||
|
||||
Returns:
|
||||
(用户称号列表, Token使用统计)
|
||||
"""
|
||||
try:
|
||||
session_id = self._make_session_id(session_id, umo)
|
||||
|
||||
logger.info(f"开始用户称号分析, session_id: {session_id}")
|
||||
return await self.user_title_analyzer.analyze_user_titles(
|
||||
messages, user_activity, umo, top_users, session_id
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f"用户称号分析失败: {e}")
|
||||
return [], TokenUsage()
|
||||
|
||||
async def analyze_golden_quotes(
|
||||
self,
|
||||
messages: list[dict],
|
||||
umo: str | None = None,
|
||||
session_id: str | None = None,
|
||||
) -> tuple[list[GoldenQuote], TokenUsage]:
|
||||
"""
|
||||
使用LLM分析群聊金句
|
||||
保持原有接口,委托给专门的GoldenQuoteAnalyzer处理
|
||||
|
||||
Args:
|
||||
messages: 群聊消息列表
|
||||
umo: 模型唯一标识符
|
||||
session_id: 会话ID (用于调试模式)
|
||||
|
||||
Returns:
|
||||
(金句列表, Token使用统计)
|
||||
"""
|
||||
try:
|
||||
session_id = self._make_session_id(session_id, umo)
|
||||
|
||||
logger.info(f"开始金句分析, session_id: {session_id}")
|
||||
return await self.golden_quote_analyzer.analyze_golden_quotes(
|
||||
messages, umo, session_id
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f"金句分析失败: {e}")
|
||||
return [], TokenUsage()
|
||||
|
||||
async def analyze_comic_storyboards(
|
||||
self,
|
||||
topics: list[dict],
|
||||
umo: str | None = None,
|
||||
session_id: str | None = None,
|
||||
persona_id: str | None = None,
|
||||
prompt_template: str | None = None,
|
||||
) -> tuple[list[dict], TokenUsage]:
|
||||
"""使用 LLM 分析并生成漫画分镜和绘画提示词。
|
||||
|
||||
Args:
|
||||
topics: 已提取的有效群聊话题。
|
||||
umo: 群聊统一消息来源标识。
|
||||
session_id: 调试会话标识。
|
||||
persona_id: 漫画分镜专用人格 ID。
|
||||
prompt_template: 角色专属的漫画分镜提示词模板。
|
||||
|
||||
Returns:
|
||||
分镜列表和 Token 使用统计。
|
||||
"""
|
||||
try:
|
||||
session_id = self._make_session_id(session_id, umo)
|
||||
|
||||
logger.info(f"开始漫画分镜分析, session_id: {session_id}")
|
||||
(
|
||||
storyboards,
|
||||
usage,
|
||||
) = await self.comic_storyboard_analyzer.analyze_storyboards(
|
||||
topics, umo, session_id, persona_id, prompt_template
|
||||
)
|
||||
trace = TraceContext.current()
|
||||
if trace and usage and usage.total_tokens > 0:
|
||||
trace.add_token_usage(
|
||||
usage.prompt_tokens,
|
||||
usage.completion_tokens,
|
||||
analyzer_name="comic_storyboard",
|
||||
)
|
||||
return storyboards, usage
|
||||
except Exception as e:
|
||||
logger.error(f"漫画分镜分析失败: {e}", exc_info=True)
|
||||
return [], TokenUsage()
|
||||
|
||||
async def summarize_quality_reviews(
|
||||
self,
|
||||
batch_reviews: list[dict],
|
||||
umo: str | None = None,
|
||||
session_id: str | None = None,
|
||||
) -> tuple[QualityReview | None, TokenUsage]:
|
||||
"""
|
||||
汇总多个质量分析报告(增量模式使用)
|
||||
"""
|
||||
return await self.chat_quality_analyzer.summarize_batch_reviews(
|
||||
batch_reviews, umo, session_id
|
||||
)
|
||||
|
||||
async def analyze_all_concurrent(
|
||||
self,
|
||||
messages: list[dict],
|
||||
user_activity: dict,
|
||||
umo: str | None = None,
|
||||
top_users: list[dict] | None = None,
|
||||
topic_enabled: bool = True,
|
||||
user_title_enabled: bool = True,
|
||||
golden_quote_enabled: bool = True,
|
||||
chat_quality_enabled: bool = False,
|
||||
) -> tuple[
|
||||
list[SummaryTopic],
|
||||
list[UserTitle],
|
||||
list[GoldenQuote],
|
||||
TokenUsage,
|
||||
QualityReview | None,
|
||||
]:
|
||||
"""
|
||||
并发执行所有分析任务(话题、用户称号、金句),支持按需启用。
|
||||
|
||||
Args:
|
||||
messages: 群聊消息列表
|
||||
user_activity: 用户分析统计
|
||||
umo: 模型唯一标识符
|
||||
top_users: 活跃用户列表(可选)
|
||||
topic_enabled: 是否启用话题分析
|
||||
user_title_enabled: 是否启用用户称号分析
|
||||
golden_quote_enabled: 是否启用金句分析
|
||||
|
||||
Returns:
|
||||
(话题列表, 用户称号列表, 金句列表, 总Token使用统计)
|
||||
"""
|
||||
try:
|
||||
session_id = self._make_session_id(None, umo)
|
||||
|
||||
logger.info(
|
||||
f"开始并发执行分析任务 (话题:{topic_enabled}, 称号:{user_title_enabled}, 金句:{golden_quote_enabled}, 质量:{chat_quality_enabled}),会话ID: {session_id}"
|
||||
)
|
||||
|
||||
# 保存原始消息数据 (Debug Mode)
|
||||
if self.config_manager.get_debug_mode():
|
||||
self._save_debug_messages(messages, session_id)
|
||||
|
||||
# 构建并发任务列表
|
||||
tasks = []
|
||||
task_names = []
|
||||
|
||||
if topic_enabled:
|
||||
tasks.append(
|
||||
self.topic_analyzer.analyze_topics(messages, umo, session_id)
|
||||
)
|
||||
task_names.append("topic")
|
||||
|
||||
if user_title_enabled:
|
||||
tasks.append(
|
||||
self.user_title_analyzer.analyze_user_titles(
|
||||
messages, user_activity, umo, top_users, session_id
|
||||
)
|
||||
)
|
||||
task_names.append("user_title")
|
||||
|
||||
if golden_quote_enabled:
|
||||
tasks.append(
|
||||
self.golden_quote_analyzer.analyze_golden_quotes(
|
||||
messages, umo, session_id
|
||||
)
|
||||
)
|
||||
task_names.append("golden_quote")
|
||||
|
||||
if chat_quality_enabled:
|
||||
tasks.append(
|
||||
self.chat_quality_analyzer.analyze_quality(
|
||||
messages, umo, session_id
|
||||
)
|
||||
)
|
||||
task_names.append("chat_quality")
|
||||
|
||||
if not tasks:
|
||||
return [], [], [], TokenUsage(), None
|
||||
|
||||
results = await asyncio.gather(*tasks, return_exceptions=True)
|
||||
|
||||
# 处理结果
|
||||
topics, topic_usage = [], TokenUsage()
|
||||
user_titles, title_usage = [], TokenUsage()
|
||||
golden_quotes, quote_usage = [], TokenUsage()
|
||||
chat_quality_review = None
|
||||
quality_usage = TokenUsage() # Initialize here
|
||||
subtask_errors: list[str] = []
|
||||
|
||||
for i, result in enumerate(results):
|
||||
name = task_names[i]
|
||||
if isinstance(result, Exception):
|
||||
logger.error(f"分析任务 {name} 失败: {result}")
|
||||
subtask_errors.append(f"{name}: {str(result)}")
|
||||
continue
|
||||
|
||||
if name == "topic" and isinstance(result, tuple):
|
||||
topics, topic_usage = result
|
||||
elif name == "user_title" and isinstance(result, tuple):
|
||||
user_titles, title_usage = result
|
||||
elif name == "golden_quote" and isinstance(result, tuple):
|
||||
golden_quotes, quote_usage = result
|
||||
elif name == "chat_quality" and isinstance(result, tuple):
|
||||
chat_quality_review, quality_usage = result
|
||||
if not isinstance(quality_usage, TokenUsage):
|
||||
quality_usage = TokenUsage()
|
||||
|
||||
# 合并Token使用统计
|
||||
total_usage = TokenUsage(
|
||||
prompt_tokens=topic_usage.prompt_tokens
|
||||
+ title_usage.prompt_tokens
|
||||
+ quote_usage.prompt_tokens
|
||||
+ quality_usage.prompt_tokens,
|
||||
completion_tokens=topic_usage.completion_tokens
|
||||
+ title_usage.completion_tokens
|
||||
+ quote_usage.completion_tokens
|
||||
+ quality_usage.completion_tokens,
|
||||
total_tokens=topic_usage.total_tokens
|
||||
+ title_usage.total_tokens
|
||||
+ quote_usage.total_tokens
|
||||
+ quality_usage.total_tokens,
|
||||
)
|
||||
|
||||
# 校验并补全未成功产出内容的子任务说明
|
||||
if (
|
||||
topic_enabled
|
||||
and not topics
|
||||
and not any(e.startswith("topic") for e in subtask_errors)
|
||||
):
|
||||
subtask_errors.append(
|
||||
"topic: 未能提取出有效话题(有效文本过少或模型未返回话题)"
|
||||
)
|
||||
if (
|
||||
user_title_enabled
|
||||
and not user_titles
|
||||
and not any(e.startswith("user_title") for e in subtask_errors)
|
||||
):
|
||||
subtask_errors.append(
|
||||
"user_title: 未能生成用户称号(活跃用户不足或模型未返回称号)"
|
||||
)
|
||||
if (
|
||||
golden_quote_enabled
|
||||
and not golden_quotes
|
||||
and not any(e.startswith("golden_quote") for e in subtask_errors)
|
||||
):
|
||||
subtask_errors.append(
|
||||
"golden_quote: 未能提取出精彩金句(符合条件的消息过少或模型未返回)"
|
||||
)
|
||||
if (
|
||||
chat_quality_enabled
|
||||
and not chat_quality_review
|
||||
and not any(e.startswith("chat_quality") for e in subtask_errors)
|
||||
):
|
||||
subtask_errors.append(
|
||||
"chat_quality: 未能生成质量锐评(模型未按预期格式输出)"
|
||||
)
|
||||
|
||||
# 记录 Token 消耗与丰富执行详情到 TraceContext
|
||||
trace = TraceContext.current()
|
||||
if trace:
|
||||
if topic_usage.total_tokens > 0:
|
||||
trace.add_token_usage(
|
||||
topic_usage.prompt_tokens,
|
||||
topic_usage.completion_tokens,
|
||||
analyzer_name="topics",
|
||||
)
|
||||
if title_usage.total_tokens > 0:
|
||||
trace.add_token_usage(
|
||||
title_usage.prompt_tokens,
|
||||
title_usage.completion_tokens,
|
||||
analyzer_name="user_titles",
|
||||
)
|
||||
if quote_usage.total_tokens > 0:
|
||||
trace.add_token_usage(
|
||||
quote_usage.prompt_tokens,
|
||||
quote_usage.completion_tokens,
|
||||
analyzer_name="golden_quotes",
|
||||
)
|
||||
if quality_usage.total_tokens > 0:
|
||||
trace.add_token_usage(
|
||||
quality_usage.prompt_tokens,
|
||||
quality_usage.completion_tokens,
|
||||
analyzer_name="chat_quality",
|
||||
)
|
||||
# 丰富 LLM_ANALYSIS span payload 便于 WebUI 详情精准诊断
|
||||
for s in reversed(trace._spans):
|
||||
if s.get("stage_name") == "LLM_ANALYSIS":
|
||||
s.setdefault("payload", {}).update(
|
||||
{
|
||||
"topics_count": len(topics),
|
||||
"topics": [t.topic for t in topics] if topics else [],
|
||||
"user_titles_count": len(user_titles),
|
||||
"golden_quotes_count": len(golden_quotes),
|
||||
"chat_quality_review": bool(chat_quality_review),
|
||||
"prompt_tokens": total_usage.prompt_tokens,
|
||||
"completion_tokens": total_usage.completion_tokens,
|
||||
"total_tokens": total_usage.total_tokens,
|
||||
"enabled_features": {
|
||||
"topics": topic_enabled,
|
||||
"user_titles": user_title_enabled,
|
||||
"golden_quotes": golden_quote_enabled,
|
||||
"chat_quality": chat_quality_enabled,
|
||||
},
|
||||
"prompts": trace.metadata.get("llm_prompts", {}),
|
||||
}
|
||||
)
|
||||
if subtask_errors:
|
||||
s["payload"]["subtask_errors"] = subtask_errors
|
||||
break
|
||||
|
||||
logger.info(
|
||||
f"并发分析完成 - 话题: {len(topics)}, 称号: {len(user_titles)}, 金句: {len(golden_quotes)}, 质量锐评: {1 if chat_quality_review else 0}"
|
||||
)
|
||||
return (
|
||||
topics,
|
||||
user_titles,
|
||||
golden_quotes,
|
||||
total_usage,
|
||||
chat_quality_review,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"并发分析失败: {e}")
|
||||
trace = TraceContext.current()
|
||||
if trace:
|
||||
for s in reversed(trace._spans):
|
||||
if s.get("stage_name") == "LLM_ANALYSIS":
|
||||
s.setdefault("payload", {}).update(
|
||||
{
|
||||
"error": str(e),
|
||||
"subtask_errors": [f"全局并发分析异常: {e}"],
|
||||
}
|
||||
)
|
||||
break
|
||||
return [], [], [], TokenUsage(), None
|
||||
|
||||
async def analyze_incremental_concurrent(
|
||||
self,
|
||||
messages: list[dict],
|
||||
umo: str | None = None,
|
||||
topics_per_batch: int = 2,
|
||||
quotes_per_batch: int = 1,
|
||||
topic_enabled: bool = True,
|
||||
golden_quote_enabled: bool = True,
|
||||
chat_quality_enabled: bool = False,
|
||||
) -> tuple[list[SummaryTopic], list[GoldenQuote], TokenUsage, QualityReview | None]:
|
||||
"""
|
||||
增量分析模式的并发执行方法。
|
||||
仅执行话题分析和金句分析(用户称号分析在最终报告时执行),
|
||||
使用较小的批次数量以控制单次分析的输出规模。
|
||||
|
||||
Args:
|
||||
messages: 本次增量分析的群聊消息列表
|
||||
umo: 模型唯一标识符
|
||||
topics_per_batch: 本次批次最大话题数量
|
||||
quotes_per_batch: 本次批次最大金句数量
|
||||
topic_enabled: 是否启用话题分析
|
||||
golden_quote_enabled: 是否启用金句分析
|
||||
|
||||
Returns:
|
||||
(话题列表, 金句列表, 总Token使用统计)
|
||||
"""
|
||||
try:
|
||||
session_id = self._make_session_id(None, umo, "incr_")
|
||||
|
||||
logger.info(
|
||||
f"开始增量并发分析 (话题:{topic_enabled}/{topics_per_batch}, 金句:{golden_quote_enabled}/{quotes_per_batch}, 质量锐评:{chat_quality_enabled}),"
|
||||
f"消息数量: {len(messages)},会话ID: {session_id}"
|
||||
)
|
||||
|
||||
# 保存原始消息数据 (Debug Mode)
|
||||
if self.config_manager.get_debug_mode():
|
||||
self._save_debug_messages(messages, session_id)
|
||||
|
||||
# 设置增量模式的最大数量覆盖值
|
||||
self.topic_analyzer._incremental_max_count = topics_per_batch
|
||||
self.golden_quote_analyzer._incremental_max_count = quotes_per_batch
|
||||
|
||||
try:
|
||||
# 构建并发任务列表(仅话题和金句,不包含用户称号)
|
||||
tasks = []
|
||||
task_names = []
|
||||
|
||||
if topic_enabled:
|
||||
tasks.append(
|
||||
self.topic_analyzer.analyze_topics(messages, umo, session_id)
|
||||
)
|
||||
task_names.append("topic")
|
||||
|
||||
if golden_quote_enabled:
|
||||
tasks.append(
|
||||
self.golden_quote_analyzer.analyze_golden_quotes(
|
||||
messages, umo, session_id
|
||||
)
|
||||
)
|
||||
task_names.append("golden_quote")
|
||||
|
||||
if chat_quality_enabled:
|
||||
tasks.append(
|
||||
self.chat_quality_analyzer.analyze_quality(
|
||||
messages, umo, session_id
|
||||
)
|
||||
)
|
||||
task_names.append("chat_quality")
|
||||
|
||||
if not tasks:
|
||||
return [], [], TokenUsage(), None
|
||||
|
||||
results = await asyncio.gather(*tasks, return_exceptions=True)
|
||||
|
||||
# 处理结果
|
||||
topics, topic_usage = [], TokenUsage()
|
||||
golden_quotes, quote_usage = [], TokenUsage()
|
||||
chat_quality_review = None
|
||||
quality_usage = TokenUsage()
|
||||
subtask_errors: list[str] = []
|
||||
|
||||
for i, result in enumerate(results):
|
||||
name = task_names[i]
|
||||
if isinstance(result, Exception):
|
||||
logger.error(f"增量{name}分析失败: {result}")
|
||||
subtask_errors.append(f"{name}: {str(result)}")
|
||||
continue
|
||||
|
||||
if name == "topic" and isinstance(result, tuple):
|
||||
topics, topic_usage = result
|
||||
elif name == "golden_quote" and isinstance(result, tuple):
|
||||
golden_quotes, quote_usage = result
|
||||
elif name == "chat_quality" and isinstance(result, tuple):
|
||||
chat_quality_review, quality_usage = result
|
||||
if not isinstance(quality_usage, TokenUsage):
|
||||
quality_usage = TokenUsage()
|
||||
|
||||
# 校验并补全增量子任务未产出说明
|
||||
if (
|
||||
topic_enabled
|
||||
and not topics
|
||||
and not any(e.startswith("topic") for e in subtask_errors)
|
||||
):
|
||||
subtask_errors.append(
|
||||
"topic: 未能提取出增量话题(可能有效文本过少或模型未返回)"
|
||||
)
|
||||
if (
|
||||
golden_quote_enabled
|
||||
and not golden_quotes
|
||||
and not any(e.startswith("golden_quote") for e in subtask_errors)
|
||||
):
|
||||
subtask_errors.append(
|
||||
"golden_quote: 未能提取出增量金句(可能符合条件的消息过少或模型未返回)"
|
||||
)
|
||||
if (
|
||||
chat_quality_enabled
|
||||
and not chat_quality_review
|
||||
and not any(e.startswith("chat_quality") for e in subtask_errors)
|
||||
):
|
||||
subtask_errors.append(
|
||||
"chat_quality: 未能生成增量质量锐评(模型未按预期格式输出)"
|
||||
)
|
||||
|
||||
# 合并Token使用统计
|
||||
total_usage = TokenUsage(
|
||||
prompt_tokens=topic_usage.prompt_tokens
|
||||
+ quote_usage.prompt_tokens
|
||||
+ quality_usage.prompt_tokens,
|
||||
completion_tokens=topic_usage.completion_tokens
|
||||
+ quote_usage.completion_tokens
|
||||
+ quality_usage.completion_tokens,
|
||||
total_tokens=topic_usage.total_tokens
|
||||
+ quote_usage.total_tokens
|
||||
+ quality_usage.total_tokens,
|
||||
)
|
||||
|
||||
# 记录 Token 消耗到 TraceContext
|
||||
trace = TraceContext.current()
|
||||
if trace:
|
||||
if topic_usage.total_tokens > 0:
|
||||
trace.add_token_usage(
|
||||
topic_usage.prompt_tokens,
|
||||
topic_usage.completion_tokens,
|
||||
analyzer_name="topics",
|
||||
)
|
||||
if quote_usage.total_tokens > 0:
|
||||
trace.add_token_usage(
|
||||
quote_usage.prompt_tokens,
|
||||
quote_usage.completion_tokens,
|
||||
analyzer_name="golden_quotes",
|
||||
)
|
||||
if quality_usage.total_tokens > 0:
|
||||
trace.add_token_usage(
|
||||
quality_usage.prompt_tokens,
|
||||
quality_usage.completion_tokens,
|
||||
analyzer_name="chat_quality",
|
||||
)
|
||||
for s in reversed(trace._spans):
|
||||
if s.get("stage_name") == "LLM_ANALYSIS":
|
||||
s.setdefault("payload", {}).update(
|
||||
{
|
||||
"incremental": True,
|
||||
"topics_count": len(topics),
|
||||
"topics": [t.topic for t in topics]
|
||||
if topics
|
||||
else [],
|
||||
"golden_quotes_count": len(golden_quotes),
|
||||
"chat_quality_review": bool(chat_quality_review),
|
||||
"prompt_tokens": total_usage.prompt_tokens,
|
||||
"completion_tokens": total_usage.completion_tokens,
|
||||
"total_tokens": total_usage.total_tokens,
|
||||
"enabled_features": {
|
||||
"topics": topic_enabled,
|
||||
"user_titles": False,
|
||||
"golden_quotes": golden_quote_enabled,
|
||||
"chat_quality": chat_quality_enabled,
|
||||
},
|
||||
}
|
||||
)
|
||||
if subtask_errors:
|
||||
s["payload"]["subtask_errors"] = subtask_errors
|
||||
break
|
||||
|
||||
logger.info(
|
||||
f"增量并发分析完成 - 话题: {len(topics)}, 金句: {len(golden_quotes)}, 质量锐评: {1 if chat_quality_review else 0}, "
|
||||
f"Token消耗: {total_usage.total_tokens}"
|
||||
)
|
||||
return topics, golden_quotes, total_usage, chat_quality_review
|
||||
|
||||
finally:
|
||||
# 无论成功或失败,都要恢复原始的最大数量设置
|
||||
self.topic_analyzer._incremental_max_count = None
|
||||
self.golden_quote_analyzer._incremental_max_count = None
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"增量并发分析失败: {e}", exc_info=True)
|
||||
trace = TraceContext.current()
|
||||
if trace:
|
||||
for s in reversed(trace._spans):
|
||||
if s.get("stage_name") == "LLM_ANALYSIS":
|
||||
s.setdefault("payload", {}).update(
|
||||
{
|
||||
"error": str(e),
|
||||
"subtask_errors": [f"全局增量并发分析异常: {e}"],
|
||||
}
|
||||
)
|
||||
break
|
||||
return [], [], TokenUsage(), None
|
||||
|
||||
def _save_debug_messages(self, messages: list[dict], session_id: str):
|
||||
"""
|
||||
保存调试消息数据到文件(Debug Mode 专用)
|
||||
|
||||
Args:
|
||||
messages: 群聊消息列表
|
||||
session_id: 会话ID
|
||||
"""
|
||||
try:
|
||||
import json
|
||||
|
||||
from ...utils.paths import get_data_dir
|
||||
|
||||
debug_dir = get_data_dir(PLUGIN_NAME) / "debug_data"
|
||||
debug_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
msg_file_path = debug_dir / f"{session_id}_messages.json"
|
||||
with open(msg_file_path, "w", encoding="utf-8") as f:
|
||||
json.dump(messages, f, ensure_ascii=False, indent=2)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# 向后兼容的方法,保持原有调用方式
|
||||
async def _call_provider_with_retry(
|
||||
self,
|
||||
provider,
|
||||
prompt: str,
|
||||
umo: str | None = None,
|
||||
provider_id_key: str | None = None,
|
||||
):
|
||||
"""
|
||||
向后兼容的LLM调用方法
|
||||
现在委托给llm_utils模块处理
|
||||
|
||||
Args:
|
||||
provider: LLM服务商实例或None(已弃用,现在使用 provider_id_key)
|
||||
prompt: 输入的提示语
|
||||
umo: 指定使用的模型唯一标识符
|
||||
provider_id_key: 配置中的 provider_id 键名(可选)
|
||||
|
||||
Returns:
|
||||
LLM生成的结果
|
||||
"""
|
||||
return await call_provider_with_retry(
|
||||
self.context,
|
||||
self.config_manager,
|
||||
prompt,
|
||||
umo,
|
||||
provider_id_key,
|
||||
observation_label=provider_id_key or "兼容LLM调用入口",
|
||||
)
|
||||
|
||||
def _fix_json(self, text: str) -> str:
|
||||
"""
|
||||
向后兼容的JSON修复方法
|
||||
现在委托给json_utils模块处理
|
||||
|
||||
Args:
|
||||
text: 需要修复的JSON文本
|
||||
|
||||
Returns:
|
||||
修复后的JSON文本
|
||||
"""
|
||||
return fix_json(text)
|
||||
|
||||
async def analyze_retry_prompt(
|
||||
self, original_prompt: str, last_error: str, umo: str | None
|
||||
) -> str | None:
|
||||
"""
|
||||
当画图 API 遇到多次失败后,将错误信息交给 LLM 进行分析和改写。
|
||||
如果 LLM 认为原 Prompt 严重违规且无法修改,将返回 None;
|
||||
否则返回脱敏/重写后的新 Prompt,进行最后一次尝试。
|
||||
"""
|
||||
prompt = f"""
|
||||
你是一个专业且注重安全合规的内容改写员。
|
||||
有一段画图提示词在提交给画图模型时被拒绝或遇到了异常,原因可能包含敏感内容审查、尺寸格式报错或连接异常。
|
||||
|
||||
【原画图提示词】:
|
||||
{original_prompt}
|
||||
|
||||
【画图模型返回的最后一次异常信息】:
|
||||
{last_error}
|
||||
|
||||
请你根据异常信息,对原画图提示词进行诊断和修改:
|
||||
1. 如果报错是因为“色情、暴力、血腥、政治”等严重违规审查,且你认为原内容**绝对无法**被修改为健康场景(例如要求本身就是极端不合法的),请直接返回 {{"can_fix": false, "new_prompt": ""}}
|
||||
2. 如果是因为审查问题,但你可以通过**去掉敏感词**、**把场景转换为正能量/健康搞笑/委婉抽象**的画面描述来避开审查,请进行脱敏重写。
|
||||
3. 如果只是普通的超时或未知错误,你可以尝试简化画面中的复杂要素,让场景更简洁。
|
||||
|
||||
请严格以 JSON 格式输出,不要包含任何 markdown 代码块(如 ```json 等),只输出 JSON 字符串:
|
||||
{{
|
||||
"can_fix": true,
|
||||
"new_prompt": "修改后且保证健康合规的全新英文或中文画图提示词"
|
||||
}}
|
||||
"""
|
||||
try:
|
||||
llm_response = await call_provider_with_retry(
|
||||
context=self.context,
|
||||
config_manager=self.config_manager,
|
||||
prompt=prompt,
|
||||
umo=umo,
|
||||
provider_id_key="drawing_prompt_provider_id",
|
||||
observation_label="绘图提示词修复",
|
||||
)
|
||||
if not llm_response or not llm_response.completion_text:
|
||||
return None
|
||||
|
||||
response_text = llm_response.completion_text.strip()
|
||||
if response_text.startswith("```"):
|
||||
import re
|
||||
|
||||
response_text = re.sub(r"^```(?:json)?\s*", "", response_text)
|
||||
response_text = re.sub(r"\s*```$", "", response_text)
|
||||
|
||||
import json
|
||||
|
||||
try:
|
||||
data = json.loads(response_text)
|
||||
except json.JSONDecodeError:
|
||||
# 尝试通过正则寻找大括号内的内容
|
||||
import re
|
||||
|
||||
match = re.search(r"(\{.*\})", response_text, re.DOTALL)
|
||||
if match:
|
||||
data = json.loads(match.group(1))
|
||||
else:
|
||||
raise
|
||||
|
||||
if data.get("can_fix") and data.get("new_prompt"):
|
||||
new_prompt = data["new_prompt"].strip()
|
||||
if new_prompt:
|
||||
logger.info("[Comic] LLM 成功分析异常并给出了重写的安全提示词。")
|
||||
return new_prompt
|
||||
|
||||
logger.info("[Comic] LLM 判断该异常无法通过重写修复,或未提供新提示词。")
|
||||
return None
|
||||
except Exception as e:
|
||||
logger.error(f"[Comic] 请求 LLM 重写提示词时发生错误: {e}")
|
||||
return None
|
||||
@@ -0,0 +1,37 @@
|
||||
"""
|
||||
分析工具模块
|
||||
包含JSON处理和LLM API请求处理工具
|
||||
"""
|
||||
|
||||
from .info_utils import InfoUtils
|
||||
from .json_utils import (
|
||||
extract_golden_quotes_with_regex,
|
||||
extract_quality_with_regex,
|
||||
extract_topics_with_regex,
|
||||
extract_user_titles_with_regex,
|
||||
fix_json,
|
||||
parse_json_object_response,
|
||||
parse_json_response,
|
||||
)
|
||||
from .llm_utils import (
|
||||
call_provider_with_retry,
|
||||
extract_response_text,
|
||||
extract_token_usage,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
# JSON processing utilities
|
||||
"fix_json",
|
||||
"parse_json_response",
|
||||
"parse_json_object_response",
|
||||
"extract_topics_with_regex",
|
||||
"extract_user_titles_with_regex",
|
||||
"extract_golden_quotes_with_regex",
|
||||
"extract_quality_with_regex",
|
||||
# LLM utilities
|
||||
"call_provider_with_retry",
|
||||
"extract_token_usage",
|
||||
"extract_response_text",
|
||||
# Info utilities
|
||||
"InfoUtils",
|
||||
]
|
||||
@@ -0,0 +1,21 @@
|
||||
class InfoUtils:
|
||||
@staticmethod
|
||||
def get_user_nickname(config_manager, sender) -> str:
|
||||
"""
|
||||
获取用户昵称
|
||||
|
||||
优先使用nickname字段,如果为空则使用card(群名片)字段
|
||||
"""
|
||||
enable_user_card = config_manager.get_enable_user_card()
|
||||
if enable_user_card:
|
||||
return (
|
||||
sender.get("card", "")
|
||||
or sender.get("nickname", "")
|
||||
or str(sender.get("user_id", ""))
|
||||
)
|
||||
else:
|
||||
return (
|
||||
sender.get("nickname", "")
|
||||
or sender.get("card", "")
|
||||
or str(sender.get("user_id", ""))
|
||||
)
|
||||
@@ -0,0 +1,360 @@
|
||||
"""
|
||||
JSON处理工具模块
|
||||
提供JSON解析、修复和正则提取功能
|
||||
"""
|
||||
|
||||
import json
|
||||
import re
|
||||
from typing import Any
|
||||
|
||||
from ....utils.logger import logger
|
||||
|
||||
|
||||
def fix_json(text: str) -> str:
|
||||
"""
|
||||
修复JSON格式问题,包括中文符号替换
|
||||
|
||||
Args:
|
||||
text: 需要修复的JSON文本
|
||||
|
||||
Returns:
|
||||
修复后的JSON文本
|
||||
"""
|
||||
try:
|
||||
# 1. 移除markdown代码块标记
|
||||
text = re.sub(r"```json\s*", "", text)
|
||||
text = re.sub(r"```\s*$", "", text)
|
||||
|
||||
# 2. 基础清理
|
||||
text = text.replace("\n", " ").replace("\r", " ")
|
||||
text = re.sub(r"\s+", " ", text)
|
||||
|
||||
# 3. 替换中文符号为英文符号(修复)
|
||||
# 中文引号 -> 英文引号
|
||||
text = text.replace("“", '"').replace("”", '"')
|
||||
text = text.replace("‘", "'").replace("’", "'")
|
||||
# 中文逗号 -> 英文逗号
|
||||
text = text.replace(",", ",")
|
||||
# 中文冒号 -> 英文冒号
|
||||
text = text.replace(":", ":")
|
||||
# 中文括号 -> 英文括号
|
||||
text = text.replace("(", "(").replace(")", ")")
|
||||
text = text.replace("【", "[").replace("】", "]")
|
||||
|
||||
# 4. 处理字符串内容中的特殊字符
|
||||
# 转义字符串内的双引号
|
||||
def escape_quotes_in_strings(match):
|
||||
content = match.group(1)
|
||||
# 转义内部的双引号
|
||||
content = content.replace('"', '\\"')
|
||||
return f'"{content}"'
|
||||
|
||||
# 先处理字段值中的引号
|
||||
text = re.sub(r'"([^"]*(?:"[^"]*)*)"', escape_quotes_in_strings, text)
|
||||
|
||||
# 5. 修复截断的JSON
|
||||
if not text.endswith("]"):
|
||||
last_complete = text.rfind("}")
|
||||
if last_complete > 0:
|
||||
text = text[: last_complete + 1] + "]"
|
||||
|
||||
# 6. 修复常见的JSON格式问题
|
||||
# 1. 修复缺失的逗号
|
||||
text = re.sub(r"}\s*{", "}, {", text)
|
||||
|
||||
# 2. 确保字段名有引号(仅在对象开始或逗号后,避免破坏字符串值)
|
||||
def quote_field_names(match):
|
||||
prefix = match.group(1)
|
||||
key = match.group(2)
|
||||
return f'{prefix}"{key}":'
|
||||
|
||||
# 只在 { 或 , 后面匹配字段名,避免在字符串值中误匹配
|
||||
text = re.sub(r"([{,]\s*)([a-zA-Z_][a-zA-Z0-9_]*)\s*:", quote_field_names, text)
|
||||
|
||||
# 3. 移除多余的逗号
|
||||
text = re.sub(r",\s*}", "}", text)
|
||||
text = re.sub(r",\s*]", "]", text)
|
||||
|
||||
return text.strip()
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"JSON修复失败: {e}")
|
||||
return text
|
||||
|
||||
|
||||
def _parse_json_with_pattern(
|
||||
result_text: str, pattern: str, data_type: str, expected_type_name: str = "数据"
|
||||
) -> tuple[bool, Any, str | None]:
|
||||
"""
|
||||
通用内部 JSON 解析逻辑,包含提取、直接解析、修复后重试。
|
||||
"""
|
||||
fixed_json_text = None
|
||||
try:
|
||||
# 1. 基础清理:去除 markdown 代码块标记
|
||||
clean_text = result_text.strip()
|
||||
clean_text = re.sub(r"```(?:json)?\s*", "", clean_text)
|
||||
clean_text = re.sub(r"```\s*$", "", clean_text)
|
||||
|
||||
# 2. 提取 JSON 部分
|
||||
json_match = re.search(pattern, clean_text, re.DOTALL)
|
||||
if not json_match:
|
||||
error_msg = f"{data_type}响应中未找到JSON{expected_type_name}"
|
||||
logger.warning(error_msg)
|
||||
return False, None, error_msg
|
||||
|
||||
json_text = json_match.group()
|
||||
logger.debug(f"{data_type}分析JSON原文: {json_text[:500]}...")
|
||||
|
||||
# 3. 尝试直接解析
|
||||
try:
|
||||
data = json.loads(json_text)
|
||||
count_info = f",包含 {len(data)} 条数据" if isinstance(data, list) else ""
|
||||
logger.info(f"{data_type}直接解析成功{count_info}")
|
||||
return True, data, None
|
||||
except json.JSONDecodeError:
|
||||
logger.debug(f"{data_type}直接解析失败,尝试修复JSON...")
|
||||
|
||||
# 4. 修复后重试
|
||||
fixed_json_text = fix_json(json_text)
|
||||
# 修复后需要重新提取,因为 fix_json 可能会改变文本结构(例如补齐括号)
|
||||
fixed_match = re.search(pattern, fixed_json_text, re.DOTALL)
|
||||
if fixed_match:
|
||||
try:
|
||||
data = json.loads(fixed_match.group())
|
||||
count_info = (
|
||||
f",包含 {len(data)} 条数据" if isinstance(data, list) else ""
|
||||
)
|
||||
logger.info(f"{data_type}修复后解析成功{count_info}")
|
||||
return True, data, None
|
||||
except json.JSONDecodeError as e:
|
||||
error_msg = f"{data_type}JSON修复后解析仍失败: {e}"
|
||||
logger.warning(error_msg)
|
||||
return False, None, error_msg
|
||||
|
||||
error_msg = f"{data_type}修复后未找到JSON{expected_type_name}"
|
||||
return False, None, error_msg
|
||||
|
||||
except Exception as e:
|
||||
error_msg = f"{data_type}解析异常: {e}"
|
||||
logger.error(error_msg)
|
||||
return False, None, error_msg
|
||||
|
||||
|
||||
def parse_json_response(
|
||||
result_text: str, data_type: str
|
||||
) -> tuple[bool, list[dict] | None, str | None]:
|
||||
"""
|
||||
统一的JSON解析方法(用于JSON数组响应)
|
||||
"""
|
||||
return _parse_json_with_pattern(
|
||||
result_text, r"\[.*\]", data_type, expected_type_name="数组"
|
||||
)
|
||||
|
||||
|
||||
def parse_json_object_response(
|
||||
result_text: str, data_type: str
|
||||
) -> tuple[bool, dict | None, str | None]:
|
||||
"""
|
||||
统一的JSON解析方法(用于JSON对象响应)
|
||||
"""
|
||||
return _parse_json_with_pattern(
|
||||
result_text, r"\{.*\}", data_type, expected_type_name="对象"
|
||||
)
|
||||
|
||||
|
||||
def _clean_json_string(text: str) -> str:
|
||||
"""
|
||||
清理 JSON 字符串中的转义字符,用于正则提取后的数据清洗。
|
||||
"""
|
||||
return text.replace('\\"', '"').replace("\\n", " ").replace("\\t", " ")
|
||||
|
||||
|
||||
def extract_topics_with_regex(result_text: str, max_topics: int) -> list[dict]:
|
||||
"""
|
||||
使用正则表达式提取话题信息
|
||||
|
||||
Args:
|
||||
result_text: 需要提取的文本
|
||||
max_topics: 最大话题数量
|
||||
|
||||
Returns:
|
||||
话题数据列表
|
||||
"""
|
||||
try:
|
||||
# 更强的正则表达式提取话题信息,处理转义字符
|
||||
# 匹配每个完整的话题对象
|
||||
topic_pattern = r'\{\s*"topic":\s*"([^"]*(?:\\.[^"]*)*)"\s*,\s*"contributors":\s*\[(.*?)\],?\s*"detail":\s*"([^"]*(?:\\.[^"]*)*)"\s*\}'
|
||||
matches = re.findall(topic_pattern, result_text, re.DOTALL)
|
||||
|
||||
if not matches:
|
||||
# 尝试更宽松的匹配
|
||||
topic_pattern = r'"topic":\s*"([^"]*(?:\\.[^"]*)*)"[^}]*"contributors":\s*\[(.*?)\][^}]*"detail":\s*"([^"]*(?:\\.[^"]*)*)"'
|
||||
matches = re.findall(topic_pattern, result_text, re.DOTALL)
|
||||
|
||||
topics = []
|
||||
for match in matches[:max_topics]:
|
||||
topic_name = match[0].strip()
|
||||
contributors_str = match[1].strip()
|
||||
detail = _clean_json_string(match[2].strip())
|
||||
|
||||
# 解析参与者列表
|
||||
contributors = [
|
||||
contrib.strip()
|
||||
for contrib in re.findall(r'"([^"]+)"', contributors_str)
|
||||
] or ["群友"]
|
||||
|
||||
topics.append(
|
||||
{
|
||||
"topic": topic_name,
|
||||
"contributors": contributors[:5], # 最多5个参与者
|
||||
"detail": detail,
|
||||
}
|
||||
)
|
||||
|
||||
logger.info(f"话题正则表达式提取成功,提取到 {len(topics)} 条有效话题内容")
|
||||
return topics
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"话题正则表达式提取失败: {e}")
|
||||
return []
|
||||
|
||||
|
||||
def extract_user_titles_with_regex(result_text: str, max_count: int) -> list[dict]:
|
||||
"""
|
||||
使用正则表达式提取用户称号信息
|
||||
|
||||
Args:
|
||||
result_text: 需要提取的文本
|
||||
max_count: 最大提取数量
|
||||
|
||||
Returns:
|
||||
用户称号数据列表
|
||||
"""
|
||||
try:
|
||||
titles = []
|
||||
|
||||
# 正则模式:匹配完整的用户称号对象
|
||||
pattern = r'\{\s*"name":\s*"([^"]*(?:\\.[^"]*)*)"\s*,\s*"user_id":\s*"([^"]+)"\s*,\s*"title":\s*"([^"]*(?:\\.[^"]*)*)"\s*,\s*"mbti":\s*"([^"]+)"\s*,\s*"reason":\s*"([^"]*(?:\\.[^"]*)*)"\s*\}'
|
||||
matches = re.findall(pattern, result_text, re.DOTALL)
|
||||
|
||||
if not matches:
|
||||
# 尝试更宽松的匹配(字段顺序可变)
|
||||
pattern = r'"name":\s*"([^"]*(?:\\.[^"]*)*)"[^}]*"user_id":\s*"([^"]+)"[^}]*"title":\s*"([^"]*(?:\\.[^"]*)*)"[^}]*"mbti":\s*"([^"]+)"[^}]*"reason":\s*"([^"]*(?:\\.[^"]*)*)"'
|
||||
matches = re.findall(pattern, result_text, re.DOTALL)
|
||||
|
||||
for match in matches[:max_count]:
|
||||
name = match[0].strip()
|
||||
user_id = match[1].strip()
|
||||
title = match[2].strip()
|
||||
mbti = match[3].strip()
|
||||
reason = _clean_json_string(match[4].strip())
|
||||
|
||||
titles.append(
|
||||
{
|
||||
"name": name,
|
||||
"user_id": user_id,
|
||||
"title": title,
|
||||
"mbti": mbti,
|
||||
"reason": reason,
|
||||
}
|
||||
)
|
||||
|
||||
logger.info(f"用户称号正则表达式提取成功,提取到 {len(titles)} 条有效用户称号")
|
||||
return titles
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"用户称号正则表达式提取失败: {e}")
|
||||
return []
|
||||
|
||||
|
||||
def extract_golden_quotes_with_regex(result_text: str, max_count: int) -> list[dict]:
|
||||
"""
|
||||
使用正则表达式提取金句信息
|
||||
|
||||
Args:
|
||||
result_text: 需要提取的文本
|
||||
max_count: 最大提取数量
|
||||
|
||||
Returns:
|
||||
金句数据列表
|
||||
"""
|
||||
try:
|
||||
quotes = []
|
||||
|
||||
# 正则模式:匹配完整的金句对象
|
||||
pattern = r'\{\s*"content":\s*"([^"]*(?:\\.[^"]*)*)"\s*,\s*"sender":\s*"([^"]*(?:\\.[^"]*)*)"\s*,\s*"reason":\s*"([^"]*(?:\\.[^"]*)*)"\s*\}'
|
||||
matches = re.findall(pattern, result_text, re.DOTALL)
|
||||
|
||||
if not matches:
|
||||
# 尝试更宽松的匹配(字段顺序可变)
|
||||
pattern = r'"content":\s*"([^"]*(?:\\.[^"]*)*)"[^}]*"sender":\s*"([^"]*(?:\\.[^"]*)*)"[^}]*"reason":\s*"([^"]*(?:\\.[^"]*)*)"'
|
||||
matches = re.findall(pattern, result_text, re.DOTALL)
|
||||
|
||||
for match in matches[:max_count]:
|
||||
content = _clean_json_string(match[0].strip())
|
||||
sender = match[1].strip()
|
||||
reason = _clean_json_string(match[2].strip())
|
||||
|
||||
quotes.append({"content": content, "sender": sender, "reason": reason})
|
||||
|
||||
logger.info(f"金句正则表达式提取成功,提取到 {len(quotes)} 条有效金句")
|
||||
return quotes
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"金句正则表达式提取失败: {e}")
|
||||
return []
|
||||
|
||||
|
||||
def extract_quality_with_regex(result_text: str) -> dict | None:
|
||||
"""
|
||||
使用正则表达式提取聊天质量分析数据
|
||||
|
||||
当 JSON 解析失败时作为降级方案使用。
|
||||
|
||||
Args:
|
||||
result_text: LLM 返回的原始文本
|
||||
|
||||
Returns:
|
||||
解析后的质量分析字典,失败返回 None
|
||||
"""
|
||||
try:
|
||||
title_m = re.search(r'"title"\s*:\s*"([^"]*(?:\\.[^"]*)*)"', result_text)
|
||||
subtitle_m = re.search(r'"subtitle"\s*:\s*"([^"]*(?:\\.[^"]*)*)"', result_text)
|
||||
summary_m = re.search(r'"summary"\s*:\s*"([^"]*(?:\\.[^"]*)*)"', result_text)
|
||||
|
||||
# Extract dimensions array
|
||||
dims_match = re.search(r'"dimensions"\s*:\s*\[(.*)\]', result_text, re.DOTALL)
|
||||
dims = []
|
||||
if dims_match:
|
||||
dim_objects = re.findall(
|
||||
r'\{[^}]*"name"\s*:\s*"([^"]*)"[^}]*'
|
||||
r'"percentage"\s*:\s*([\d.]+)[^}]*'
|
||||
r'"comment"\s*:\s*"([^"]*(?:\\.[^"]*)*)"[^}]*\}',
|
||||
dims_match.group(1),
|
||||
)
|
||||
for dm in dim_objects:
|
||||
dims.append(
|
||||
{
|
||||
"name": dm[0],
|
||||
"percentage": float(dm[1]),
|
||||
"comment": dm[2],
|
||||
}
|
||||
)
|
||||
|
||||
if not dims:
|
||||
logger.warning("聊天质量正则提取未找到有效维度数据")
|
||||
return None
|
||||
|
||||
data = {
|
||||
"title": title_m.group(1) if title_m else "聊天质量锐评",
|
||||
"subtitle": subtitle_m.group(1) if subtitle_m else "今天的群里发生了什么?",
|
||||
"dimensions": dims,
|
||||
"summary": summary_m.group(1) if summary_m else "今天也是充满活力的一天。",
|
||||
}
|
||||
|
||||
logger.info(f"聊天质量正则表达式提取成功,提取到 {len(dims)} 个维度")
|
||||
return data
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"聊天质量正则表达式提取失败: {e}")
|
||||
return None
|
||||
@@ -0,0 +1,243 @@
|
||||
"""
|
||||
LLM API 请求处理工具模块(NoneBot 版)
|
||||
|
||||
原 AstrBot 版依赖 context.get_provider_by_id / provider.text_chat_stream。
|
||||
此文件改为直接调用 OpenAI 兼容 /chat/completions 接口:
|
||||
- 端点与密钥来自 ConfigManager(llm.llm_api_base / llm.llm_api_key / llm.llm_model)
|
||||
- 保留与原插件一致的外层函数名,供分析器 / 报告生成器复用
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import random
|
||||
import time
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
|
||||
from ....shared.trace_context import TraceContext
|
||||
from ....utils.logger import logger
|
||||
from ...config.config_manager import ConfigManager
|
||||
from .structured_output_schema import JSONObject, JSONValue
|
||||
|
||||
|
||||
@dataclass
|
||||
class LLMResponse:
|
||||
"""轻量 LLM 响应对象,替代原 AstrBot 的 astrbot.api.provider.LLMResponse。"""
|
||||
|
||||
role: str = "assistant"
|
||||
completion_text: str = ""
|
||||
usage: Any = None
|
||||
raw_completion: Any = None
|
||||
is_chunk: bool = False
|
||||
|
||||
def __str__(self) -> str:
|
||||
return self.completion_text
|
||||
|
||||
|
||||
def _build_chat_url(base_url: str) -> str:
|
||||
base = (base_url or "").strip().rstrip("/")
|
||||
if not base:
|
||||
return ""
|
||||
if base.endswith("/chat/completions"):
|
||||
return base
|
||||
return f"{base}/chat/completions"
|
||||
|
||||
|
||||
def _merge_generate_kwargs(base: dict, extra: dict | None) -> dict:
|
||||
if not extra:
|
||||
return base
|
||||
merged = dict(base)
|
||||
for k, v in extra.items():
|
||||
if k in {"stream", "stream_options"}:
|
||||
continue
|
||||
merged[k] = v
|
||||
return merged
|
||||
|
||||
|
||||
async def _call_openai(
|
||||
config_manager: ConfigManager,
|
||||
prompt: str,
|
||||
system_prompt: str | None,
|
||||
response_format: JSONObject | None,
|
||||
extra_generate_kwargs: dict[str, JSONValue] | None,
|
||||
allow_response_format: bool = True,
|
||||
) -> LLMResponse:
|
||||
base_url = config_manager.get_llm_api_base()
|
||||
api_key = config_manager.get_llm_api_key()
|
||||
model = config_manager.get_llm_model()
|
||||
timeout = config_manager.get_llm_timeout()
|
||||
if not base_url or not api_key or not model:
|
||||
raise RuntimeError("LLM 配置不完整:需要 llm_api_base / llm_api_key / llm_model")
|
||||
|
||||
url = _build_chat_url(base_url)
|
||||
if not url:
|
||||
raise RuntimeError("LLM API Base 为空,无法构造 /chat/completions 地址")
|
||||
|
||||
payload: dict[str, Any] = {"model": model, "messages": []}
|
||||
if system_prompt:
|
||||
payload["messages"].append({"role": "system", "content": system_prompt})
|
||||
payload["messages"].append({"role": "user", "content": prompt})
|
||||
payload = _merge_generate_kwargs(payload, extra_generate_kwargs)
|
||||
if response_format is not None and allow_response_format:
|
||||
payload.setdefault("response_format", {"type": "json_object"})
|
||||
|
||||
headers = {
|
||||
"Authorization": f"Bearer {api_key}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
|
||||
async with httpx.AsyncClient(timeout=timeout) as client:
|
||||
resp = await client.post(url, json=payload, headers=headers)
|
||||
if resp.status_code >= 400:
|
||||
raise RuntimeError(
|
||||
f"LLM HTTP {resp.status_code}: {resp.text[:500]}"
|
||||
)
|
||||
data = resp.json()
|
||||
try:
|
||||
content = data["choices"][0]["message"]["content"] or ""
|
||||
except (KeyError, IndexError, TypeError) as e:
|
||||
raise RuntimeError(f"LLM 返回结构异常: {e}; {str(data)[:500]}") from e
|
||||
usage = data.get("usage")
|
||||
return LLMResponse(
|
||||
role="assistant",
|
||||
completion_text=str(content),
|
||||
usage=usage,
|
||||
raw_completion=data,
|
||||
)
|
||||
|
||||
|
||||
async def get_provider_id_with_fallback(
|
||||
context: Any,
|
||||
config_manager: ConfigManager,
|
||||
provider_id_key: str | None = None,
|
||||
umo: str | None = None,
|
||||
) -> str | None:
|
||||
"""原插件用于在 AstrBot 中选择 Provider;NoneBot 版仅返回配置中的模型标识作为日志标签。"""
|
||||
pid = config_manager.get_llm_provider_id() or config_manager.get_llm_model() or "default"
|
||||
logger.debug(f"[Provider 选择] NoneBot 统一 LLM 通道, label={pid}")
|
||||
return pid
|
||||
|
||||
|
||||
async def call_provider_with_retry(
|
||||
context: Any,
|
||||
config_manager: ConfigManager,
|
||||
prompt: str,
|
||||
umo: str | None = None,
|
||||
provider_id_key: str | None = None,
|
||||
provider_id: str | None = None,
|
||||
system_prompt: str | None = None,
|
||||
response_format: JSONObject | None = None,
|
||||
extra_generate_kwargs: dict[str, JSONValue] | None = None,
|
||||
observation_label: str | None = None,
|
||||
) -> LLMResponse | None:
|
||||
"""调用 OpenAI 兼容端点,带重试与退避。失败返回 None。"""
|
||||
if not prompt or not prompt.strip():
|
||||
logger.error("LLM provider: prompt 为空,无法调用")
|
||||
return None
|
||||
|
||||
retries = max(1, int(config_manager.get_llm_retries()))
|
||||
backoff = max(0, float(config_manager.get_llm_backoff()))
|
||||
trace = TraceContext.current()
|
||||
area = observation_label or provider_id_key or "未标注"
|
||||
|
||||
last_exc: Exception | None = None
|
||||
current_response_format = response_format
|
||||
for attempt in range(1, retries + 1):
|
||||
try:
|
||||
resp = await _call_openai(
|
||||
config_manager,
|
||||
prompt,
|
||||
system_prompt,
|
||||
current_response_format,
|
||||
extra_generate_kwargs,
|
||||
allow_response_format=current_response_format is not None,
|
||||
)
|
||||
if trace:
|
||||
trace.metadata.setdefault("llm_attempts", []).append(
|
||||
{
|
||||
"area": area,
|
||||
"attempt": attempt,
|
||||
"provider_id": provider_id or "default",
|
||||
"status": "success",
|
||||
}
|
||||
)
|
||||
return resp
|
||||
except Exception as e:
|
||||
last_exc = e
|
||||
logger.warning(
|
||||
f"[LLM 调用] {area} 第 {attempt}/{retries} 次请求失败: {e}"
|
||||
)
|
||||
# response_format 不被支持时,降级去掉 schema 再试一次
|
||||
if (
|
||||
current_response_format is not None
|
||||
and "response_format" in str(e).lower()
|
||||
):
|
||||
logger.warning(f"[LLM 调用] {area} 关闭 response_format 重试")
|
||||
current_response_format = None
|
||||
try:
|
||||
resp = await _call_openai(
|
||||
config_manager,
|
||||
prompt,
|
||||
system_prompt,
|
||||
None,
|
||||
extra_generate_kwargs,
|
||||
allow_response_format=False,
|
||||
)
|
||||
if trace:
|
||||
trace.metadata.setdefault("llm_attempts", []).append(
|
||||
{
|
||||
"area": area,
|
||||
"attempt": attempt,
|
||||
"provider_id": provider_id or "default",
|
||||
"status": "success",
|
||||
"note": "without_response_format",
|
||||
}
|
||||
)
|
||||
return resp
|
||||
except Exception as inner_e:
|
||||
last_exc = inner_e
|
||||
|
||||
if attempt < retries:
|
||||
sleep_time = backoff * (2 ** (attempt - 1)) + random.uniform(0, 1)
|
||||
await asyncio.sleep(sleep_time)
|
||||
|
||||
logger.error(f"[LLM 调用] {area} 重试耗尽,最终失败: {last_exc}")
|
||||
return None
|
||||
|
||||
|
||||
def extract_token_usage(response: Any) -> dict[str, int]:
|
||||
"""从 LLM 响应中提取 token 使用统计。"""
|
||||
token_usage = {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0}
|
||||
try:
|
||||
usage = getattr(response, "usage", None)
|
||||
if not usage and hasattr(response, "raw_completion"):
|
||||
raw = response.raw_completion
|
||||
usage = raw.get("usage") if isinstance(raw, dict) else getattr(raw, "usage", None)
|
||||
if usage:
|
||||
if isinstance(usage, dict):
|
||||
token_usage["prompt_tokens"] = usage.get("prompt_tokens", 0) or 0
|
||||
token_usage["completion_tokens"] = usage.get("completion_tokens", 0) or 0
|
||||
token_usage["total_tokens"] = usage.get("total_tokens", 0) or 0
|
||||
else:
|
||||
token_usage["prompt_tokens"] = getattr(usage, "prompt_tokens", 0) or 0
|
||||
token_usage["completion_tokens"] = getattr(usage, "completion_tokens", 0) or 0
|
||||
token_usage["total_tokens"] = getattr(usage, "total_tokens", 0) or 0
|
||||
return token_usage
|
||||
except Exception as e:
|
||||
logger.error(f"提取 token 使用统计失败: {e}")
|
||||
return {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0}
|
||||
|
||||
|
||||
def extract_response_text(response: Any) -> str:
|
||||
"""从 LLM 响应中提取文本内容。"""
|
||||
try:
|
||||
if hasattr(response, "completion_text"):
|
||||
return response.completion_text
|
||||
return str(response)
|
||||
except Exception as e:
|
||||
logger.error(f"提取响应文本失败: {e}")
|
||||
return ""
|
||||
@@ -0,0 +1,127 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, ValidationError, field_validator
|
||||
|
||||
|
||||
class TopicItemModel(BaseModel):
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
|
||||
topic: str
|
||||
contributors: list[str]
|
||||
detail: str
|
||||
|
||||
@field_validator("topic", "detail", mode="before")
|
||||
@classmethod
|
||||
def _normalize_text(cls, value: object) -> str:
|
||||
return str(value).strip()
|
||||
|
||||
@field_validator("contributors", mode="before")
|
||||
@classmethod
|
||||
def _normalize_contributors(cls, value: object) -> list[str]:
|
||||
if not isinstance(value, list):
|
||||
return []
|
||||
contributors: list[str] = []
|
||||
for item in value:
|
||||
text = str(item).strip()
|
||||
if text:
|
||||
contributors.append(text)
|
||||
return contributors
|
||||
|
||||
|
||||
class UserTitleItemModel(BaseModel):
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
|
||||
name: str
|
||||
user_id: str
|
||||
title: str
|
||||
mbti: str
|
||||
reason: str
|
||||
|
||||
@field_validator("name", "user_id", "title", "mbti", "reason", mode="before")
|
||||
@classmethod
|
||||
def _normalize_text(cls, value: object) -> str:
|
||||
return str(value).strip()
|
||||
|
||||
|
||||
class GoldenQuoteItemModel(BaseModel):
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
|
||||
content: str
|
||||
sender: str
|
||||
reason: str
|
||||
|
||||
@field_validator("content", "sender", "reason", mode="before")
|
||||
@classmethod
|
||||
def _normalize_text(cls, value: object) -> str:
|
||||
return str(value).strip()
|
||||
|
||||
|
||||
class QualityDimensionModel(BaseModel):
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
|
||||
name: str
|
||||
percentage: float
|
||||
comment: str
|
||||
|
||||
@field_validator("name", "comment", mode="before")
|
||||
@classmethod
|
||||
def _normalize_text(cls, value: object) -> str:
|
||||
return str(value).strip()
|
||||
|
||||
|
||||
class QualityReviewModel(BaseModel):
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
|
||||
title: str
|
||||
subtitle: str
|
||||
dimensions: list[QualityDimensionModel]
|
||||
summary: str
|
||||
|
||||
@field_validator("title", "subtitle", "summary", mode="before")
|
||||
@classmethod
|
||||
def _normalize_text(cls, value: object) -> str:
|
||||
return str(value).strip()
|
||||
|
||||
|
||||
def validate_topic_items(
|
||||
data_list: list[dict],
|
||||
) -> tuple[bool, list[dict] | None, str | None]:
|
||||
try:
|
||||
normalized = [
|
||||
TopicItemModel.model_validate(item).model_dump() for item in data_list
|
||||
]
|
||||
return True, normalized, None
|
||||
except ValidationError as e:
|
||||
return False, None, str(e)
|
||||
|
||||
|
||||
def validate_user_title_items(
|
||||
data_list: list[dict],
|
||||
) -> tuple[bool, list[dict] | None, str | None]:
|
||||
try:
|
||||
normalized = [
|
||||
UserTitleItemModel.model_validate(item).model_dump() for item in data_list
|
||||
]
|
||||
return True, normalized, None
|
||||
except ValidationError as e:
|
||||
return False, None, str(e)
|
||||
|
||||
|
||||
def validate_golden_quote_items(
|
||||
data_list: list[dict],
|
||||
) -> tuple[bool, list[dict] | None, str | None]:
|
||||
try:
|
||||
normalized = [
|
||||
GoldenQuoteItemModel.model_validate(item).model_dump() for item in data_list
|
||||
]
|
||||
return True, normalized, None
|
||||
except ValidationError as e:
|
||||
return False, None, str(e)
|
||||
|
||||
|
||||
def validate_quality_review_item(data: dict) -> tuple[bool, dict | None, str | None]:
|
||||
try:
|
||||
normalized = QualityReviewModel.model_validate(data).model_dump()
|
||||
return True, normalized, None
|
||||
except ValidationError as e:
|
||||
return False, None, str(e)
|
||||
@@ -0,0 +1,101 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TypeAlias
|
||||
|
||||
JSONScalar: TypeAlias = str | int | float | bool | None
|
||||
JSONValue: TypeAlias = JSONScalar | dict[str, "JSONValue"] | list["JSONValue"]
|
||||
JSONObject: TypeAlias = dict[str, JSONValue]
|
||||
|
||||
|
||||
def build_response_format(name: str, schema: JSONObject) -> JSONObject:
|
||||
return {
|
||||
"type": "json_schema",
|
||||
"json_schema": {
|
||||
"name": name,
|
||||
"strict": True,
|
||||
"schema": schema,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def build_topics_schema(max_items: int) -> JSONObject:
|
||||
return {
|
||||
"type": "array",
|
||||
"maxItems": max(1, int(max_items)),
|
||||
"items": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"topic": {"type": "string"},
|
||||
"contributors": {
|
||||
"type": "array",
|
||||
"items": {"type": "string"},
|
||||
},
|
||||
"detail": {"type": "string"},
|
||||
},
|
||||
"required": ["topic", "contributors", "detail"],
|
||||
"additionalProperties": False,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def build_user_titles_schema(max_items: int) -> JSONObject:
|
||||
return {
|
||||
"type": "array",
|
||||
"maxItems": max(1, int(max_items)),
|
||||
"items": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"name": {"type": "string"},
|
||||
"user_id": {"type": "string"},
|
||||
"title": {"type": "string"},
|
||||
"mbti": {"type": "string"},
|
||||
"reason": {"type": "string"},
|
||||
},
|
||||
"required": ["name", "user_id", "title", "mbti", "reason"],
|
||||
"additionalProperties": False,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def build_golden_quotes_schema(max_items: int) -> JSONObject:
|
||||
return {
|
||||
"type": "array",
|
||||
"maxItems": max(1, int(max_items)),
|
||||
"items": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"content": {"type": "string"},
|
||||
"sender": {"type": "string"},
|
||||
"reason": {"type": "string"},
|
||||
},
|
||||
"required": ["content", "sender", "reason"],
|
||||
"additionalProperties": False,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def build_chat_quality_schema(max_dimensions: int) -> JSONObject:
|
||||
return {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"title": {"type": "string"},
|
||||
"subtitle": {"type": "string"},
|
||||
"dimensions": {
|
||||
"type": "array",
|
||||
"maxItems": max(1, int(max_dimensions)),
|
||||
"items": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"name": {"type": "string"},
|
||||
"percentage": {"type": "number"},
|
||||
"comment": {"type": "string"},
|
||||
},
|
||||
"required": ["name", "percentage", "comment"],
|
||||
"additionalProperties": False,
|
||||
},
|
||||
},
|
||||
"summary": {"type": "string"},
|
||||
},
|
||||
"required": ["title", "subtitle", "dimensions", "summary"],
|
||||
"additionalProperties": False,
|
||||
}
|
||||
@@ -0,0 +1,7 @@
|
||||
"""
|
||||
Config Module - Configuration management
|
||||
"""
|
||||
|
||||
from .config_manager import ConfigManager
|
||||
|
||||
__all__ = ["ConfigManager"]
|
||||
@@ -0,0 +1,3 @@
|
||||
from .service import DrawingApiRequestService
|
||||
|
||||
__all__ = ["DrawingApiRequestService"]
|
||||
@@ -0,0 +1,89 @@
|
||||
import base64
|
||||
from math import gcd
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
|
||||
from ....utils.logger import logger
|
||||
from .context import DrawingRequestContext
|
||||
|
||||
|
||||
async def call_chat_api(
|
||||
context: DrawingRequestContext,
|
||||
prompt: str,
|
||||
images_data: list[tuple[bytes, str]] | None = None,
|
||||
provider: dict | None = None,
|
||||
) -> bytes | None:
|
||||
provider = provider or {}
|
||||
raw_url = context.get_provider_value("api_url", provider)
|
||||
target_url = context.build_target_url(raw_url, "chat")
|
||||
|
||||
api_key = context.get_provider_value("api_key", provider)
|
||||
model = context.get_provider_value("model", provider)
|
||||
|
||||
timeout = context.get_provider_value("timeout", provider)
|
||||
|
||||
headers = {
|
||||
"Authorization": f"Bearer {api_key}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
|
||||
raw_size = context.get_provider_value("image_size", provider)
|
||||
ar = context.get_provider_value("aspect_ratio", provider)
|
||||
resolved_size = context.resolve_size(raw_size, ar)
|
||||
|
||||
# 将长宽比与分辨率要求显式追加到 prompt 结尾,防止 Chat 协议模型忽略
|
||||
width, height = map(int, resolved_size.split("x", 1))
|
||||
divisor = gcd(width, height)
|
||||
effective_aspect_ratio = f"{width // divisor}:{height // divisor}"
|
||||
if width > height:
|
||||
orientation = "Horizontal Landscape Orientation"
|
||||
elif width < height:
|
||||
orientation = "Vertical Portrait Orientation"
|
||||
else:
|
||||
orientation = "Square Orientation"
|
||||
full_prompt = f"{prompt}\n\n[Image Layout & Spec Requirements: Aspect Ratio {effective_aspect_ratio}, Resolution {resolved_size}, {orientation}]"
|
||||
|
||||
content = []
|
||||
for img_bytes, mime in images_data or []:
|
||||
b64 = base64.b64encode(img_bytes).decode("utf-8")
|
||||
content.append(
|
||||
{"type": "image_url", "image_url": {"url": f"data:{mime};base64,{b64}"}}
|
||||
)
|
||||
|
||||
content.append({"type": "text", "text": full_prompt})
|
||||
|
||||
payload: dict[str, Any] = {
|
||||
"model": model,
|
||||
"messages": [{"role": "user", "content": content}],
|
||||
}
|
||||
|
||||
logger.info(
|
||||
f"[Comic] 发起 Chat API 请求 -> {context.sanitize_url(target_url)} (model={model}, size={resolved_size}, aspect_ratio={ar})..."
|
||||
)
|
||||
|
||||
api_timeout = httpx.Timeout(connect=20.0, read=timeout, write=20.0, pool=20.0)
|
||||
async with httpx.AsyncClient(
|
||||
timeout=api_timeout, proxy=context.get_request_proxy(provider)
|
||||
) as client:
|
||||
resp = await client.post(target_url, headers=headers, json=payload)
|
||||
|
||||
if resp.status_code != 200:
|
||||
snippet = resp.text[:500] if resp.text else "(空响应)"
|
||||
raise Exception(f"API 请求失败 [HTTP {resp.status_code}]: {snippet}")
|
||||
|
||||
try:
|
||||
data = resp.json()
|
||||
except Exception:
|
||||
snippet = resp.text[:500] if resp.text else "(空正文)"
|
||||
raise Exception(
|
||||
f"API 未返回合法的 JSON [HTTP {resp.status_code}]: {snippet}"
|
||||
)
|
||||
|
||||
image = await context.extract_image(data, context.get_request_proxy(provider))
|
||||
if image:
|
||||
return image
|
||||
|
||||
raise Exception(
|
||||
f"无法从 Chat API 的回复中提取到图片: {context.summarize_response(data)}"
|
||||
)
|
||||
@@ -0,0 +1,42 @@
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
|
||||
from ....utils.logger import logger
|
||||
from .context import DrawingRequestContext
|
||||
|
||||
|
||||
async def post_json_for_image(
|
||||
context: DrawingRequestContext,
|
||||
target_url: str,
|
||||
headers: dict[str, str],
|
||||
payload: dict[str, Any],
|
||||
timeout: int | float,
|
||||
provider_name: str,
|
||||
provider: dict,
|
||||
) -> bytes | None:
|
||||
"""发送 JSON 图片生成请求,并从响应中提取图片。"""
|
||||
headers["Content-Type"] = "application/json"
|
||||
api_timeout = httpx.Timeout(connect=20.0, read=timeout, write=20.0, pool=20.0)
|
||||
logger.info(
|
||||
f"[Comic] 发起 {provider_name} 图片请求 -> {context.sanitize_url(target_url)}"
|
||||
)
|
||||
async with httpx.AsyncClient(
|
||||
timeout=api_timeout, proxy=context.get_request_proxy(provider)
|
||||
) as client:
|
||||
response = await client.post(target_url, headers=headers, json=payload)
|
||||
if not 200 <= response.status_code < 300:
|
||||
message = response.text[:500] if response.text else "(空响应)"
|
||||
raise Exception(
|
||||
f"{provider_name} API 请求失败 [HTTP {response.status_code}]: {message}"
|
||||
)
|
||||
try:
|
||||
data = response.json()
|
||||
except ValueError as exc:
|
||||
raise Exception(f"{provider_name} API 未返回合法 JSON") from exc
|
||||
image = await context.extract_image(data, context.get_request_proxy(provider))
|
||||
if image:
|
||||
return image
|
||||
raise Exception(
|
||||
f"{provider_name} API 返回格式异常: {context.summarize_response(data)}"
|
||||
)
|
||||
@@ -0,0 +1,73 @@
|
||||
"""绘图请求服务与客户端之间的显式依赖契约。
|
||||
|
||||
供应商请求模块不直接依赖 ``DrawingClient``,只通过本上下文访问 URL 解析、
|
||||
配置回退、代理、尺寸换算和响应解析能力。这样服务商文件保持可独立阅读,且
|
||||
不会因继承关系隐式获得不相关的客户端状态。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Awaitable, Callable
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Protocol
|
||||
|
||||
|
||||
class DrawingRequestHooks(Protocol):
|
||||
"""描述绘图请求服务依赖的宿主能力。"""
|
||||
|
||||
def _build_target_url(self, raw_url: str, protocol: str) -> str: ...
|
||||
|
||||
def _get_provider_value(self, name: str, provider: dict) -> Any: ...
|
||||
|
||||
def _get_request_proxy(self, provider: dict | None = None) -> str | None: ...
|
||||
|
||||
def _resolve_size(self, size_or_ratio: str, aspect_ratio: str) -> str: ...
|
||||
|
||||
def _sanitize_url(self, url: str) -> str: ...
|
||||
|
||||
def _summarize_response(self, data: Any) -> str: ...
|
||||
|
||||
def _decode_base64(self, encoded: str) -> bytes: ...
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class DrawingRequestContext:
|
||||
"""聚合绘图请求执行所需的显式依赖。
|
||||
|
||||
该数据对象是组合关系的边界:请求模块只接收这个最小能力集合,HTTP JSON
|
||||
请求和图片提取则以函数形式注入,方便保留原有可替换入口。
|
||||
"""
|
||||
|
||||
hooks: DrawingRequestHooks
|
||||
request_json: Callable[..., Awaitable[bytes | None]]
|
||||
extract_image: Callable[[Any, str | None], Awaitable[bytes | None]]
|
||||
|
||||
def build_target_url(self, raw_url: str, protocol: str) -> str:
|
||||
return self.hooks._build_target_url(raw_url, protocol)
|
||||
|
||||
def get_provider_value(self, name: str, provider: dict) -> Any:
|
||||
return self.hooks._get_provider_value(name, provider)
|
||||
|
||||
def get_request_proxy(self, provider: dict | None = None) -> str | None:
|
||||
return self.hooks._get_request_proxy(provider)
|
||||
|
||||
def resolve_size(self, size_or_ratio: str, aspect_ratio: str) -> str:
|
||||
"""按当前供应商条目的宽高比解析尺寸别名。
|
||||
|
||||
Args:
|
||||
size_or_ratio: 条目中的尺寸别名、分辨率或比例。
|
||||
aspect_ratio: 同一条目中的目标宽高比。
|
||||
|
||||
Returns:
|
||||
对应的 ``宽x高`` 尺寸字符串。
|
||||
"""
|
||||
return self.hooks._resolve_size(size_or_ratio, aspect_ratio)
|
||||
|
||||
def sanitize_url(self, url: str) -> str:
|
||||
return self.hooks._sanitize_url(url)
|
||||
|
||||
def summarize_response(self, data: Any) -> str:
|
||||
return self.hooks._summarize_response(data)
|
||||
|
||||
def decode_base64(self, encoded: str) -> bytes:
|
||||
return self.hooks._decode_base64(encoded)
|
||||
@@ -0,0 +1,154 @@
|
||||
import base64
|
||||
import binascii
|
||||
import re
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
|
||||
from ....utils.logger import logger
|
||||
from .context import DrawingRequestContext
|
||||
|
||||
|
||||
async def call_gemini_api(
|
||||
context: DrawingRequestContext,
|
||||
prompt: str,
|
||||
images_data: list[tuple[bytes, str]] | None = None,
|
||||
provider: dict | None = None,
|
||||
) -> bytes | None:
|
||||
"""调用 Google Gemini Interactions 图片接口。
|
||||
|
||||
Args:
|
||||
prompt: 图片生成或编辑提示词。
|
||||
images_data: 可选参考图片及其 MIME 类型列表,当前只使用第一张。
|
||||
|
||||
Returns:
|
||||
最后一个模型输出中的图片二进制数据。
|
||||
|
||||
Raises:
|
||||
Exception: 请求失败、响应不是 JSON 或响应中没有最终图片。
|
||||
"""
|
||||
provider = provider or {}
|
||||
raw_url = context.get_provider_value("api_url", provider)
|
||||
target_url = context.build_target_url(raw_url, "gemini")
|
||||
api_key = context.get_provider_value("api_key", provider)
|
||||
model = context.get_provider_value("model", provider)
|
||||
timeout = context.get_provider_value("timeout", provider)
|
||||
aspect_ratio = context.get_provider_value("aspect_ratio", provider)
|
||||
|
||||
raw_size = str(context.get_provider_value("image_size", provider)).strip()
|
||||
if raw_size.upper() in {"1K", "2K", "4K"}:
|
||||
image_size = raw_size.upper()
|
||||
elif re.fullmatch(r"\d+x\d+", raw_size.lower()):
|
||||
width, height = map(int, raw_size.lower().split("x", 1))
|
||||
longest_edge = max(width, height)
|
||||
if longest_edge <= 1024:
|
||||
image_size = "1K"
|
||||
elif longest_edge <= 2048:
|
||||
image_size = "2K"
|
||||
else:
|
||||
image_size = "4K"
|
||||
else:
|
||||
image_size = "1K"
|
||||
|
||||
output_format = str(context.get_provider_value("output_format", provider)).lower()
|
||||
output_mime = {
|
||||
"png": "image/png",
|
||||
"jpeg": "image/jpeg",
|
||||
"jpg": "image/jpeg",
|
||||
}.get(output_format)
|
||||
|
||||
input_content: list[dict[str, str]] = [{"type": "text", "text": prompt}]
|
||||
reference_bytes = 0
|
||||
for image_bytes, mime in images_data or []:
|
||||
image_mime = mime if mime.startswith("image/") else "image/png"
|
||||
input_content.append(
|
||||
{
|
||||
"type": "image",
|
||||
"data": base64.b64encode(image_bytes).decode("ascii"),
|
||||
"mime_type": image_mime,
|
||||
}
|
||||
)
|
||||
reference_bytes += len(image_bytes)
|
||||
|
||||
response_format: dict[str, str] = {
|
||||
"type": "image",
|
||||
"aspect_ratio": aspect_ratio,
|
||||
"image_size": image_size,
|
||||
}
|
||||
if output_mime:
|
||||
response_format["mime_type"] = output_mime
|
||||
|
||||
payload: dict[str, Any] = {
|
||||
"model": model,
|
||||
"input": input_content,
|
||||
"response_format": response_format,
|
||||
"store": False,
|
||||
}
|
||||
headers = {
|
||||
"x-goog-api-key": api_key,
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
|
||||
logger.info(
|
||||
f"[Comic] 发起 Gemini Interactions API 请求 -> {context.sanitize_url(target_url)} "
|
||||
f"(model={model}, image_size={image_size}, "
|
||||
f"aspect_ratio={aspect_ratio}, reference_bytes={reference_bytes})..."
|
||||
)
|
||||
api_timeout = httpx.Timeout(connect=20.0, read=timeout, write=20.0, pool=20.0)
|
||||
async with httpx.AsyncClient(
|
||||
timeout=api_timeout, proxy=context.get_request_proxy(provider)
|
||||
) as client:
|
||||
resp = await client.post(target_url, headers=headers, json=payload)
|
||||
|
||||
if not 200 <= resp.status_code < 300:
|
||||
error_summary = resp.text[:500] if resp.text else "(空响应)"
|
||||
raise Exception(
|
||||
f"Gemini API 请求失败 [HTTP {resp.status_code}]: {error_summary}"
|
||||
)
|
||||
|
||||
try:
|
||||
data = resp.json()
|
||||
except Exception:
|
||||
raise Exception(
|
||||
f"Gemini API 未返回合法的 JSON [HTTP {resp.status_code}]: "
|
||||
f"<body len={len(resp.content)}>"
|
||||
)
|
||||
|
||||
steps = data.get("steps") if isinstance(data, dict) else None
|
||||
model_outputs = (
|
||||
[
|
||||
step
|
||||
for step in steps
|
||||
if isinstance(step, dict) and step.get("type") == "model_output"
|
||||
]
|
||||
if isinstance(steps, list)
|
||||
else []
|
||||
)
|
||||
for step in reversed(model_outputs):
|
||||
content = step.get("content")
|
||||
if not isinstance(content, list):
|
||||
continue
|
||||
for item in reversed(content):
|
||||
if not isinstance(item, dict) or item.get("type") != "image":
|
||||
continue
|
||||
encoded = item.get("data")
|
||||
if not isinstance(encoded, str) or not encoded.strip():
|
||||
continue
|
||||
try:
|
||||
return context.decode_base64(encoded)
|
||||
except (ValueError, TypeError, binascii.Error) as exc:
|
||||
logger.debug(f"[Comic] 跳过无效 Gemini 最终图片: {exc}")
|
||||
|
||||
# 当响应包含 steps 时,只在最终模型输出中回退提取图片,避免误取中间推理图。
|
||||
fallback_data: Any = model_outputs if isinstance(steps, list) else data
|
||||
image = await context.extract_image(
|
||||
fallback_data, context.get_request_proxy(provider)
|
||||
)
|
||||
if image:
|
||||
return image
|
||||
|
||||
status = data.get("status") if isinstance(data, dict) else None
|
||||
raise Exception(
|
||||
f"Gemini API 未返回最终图片 (status={status or 'unknown'}): "
|
||||
f"{context.summarize_response(data)}"
|
||||
)
|
||||
@@ -0,0 +1,53 @@
|
||||
import base64
|
||||
from typing import Any
|
||||
|
||||
from .context import DrawingRequestContext
|
||||
|
||||
|
||||
async def call_google_api(
|
||||
context: DrawingRequestContext,
|
||||
prompt: str,
|
||||
images_data: list[tuple[bytes, str]] | None,
|
||||
provider: dict,
|
||||
) -> bytes | None:
|
||||
"""调用 Google Gemini generateContent 官方接口。"""
|
||||
api_base = str(context.get_provider_value("api_url", provider)).rstrip("/")
|
||||
model = context.get_provider_value("model", provider)
|
||||
if ":generateContent" in api_base:
|
||||
target_url = api_base
|
||||
else:
|
||||
if not api_base:
|
||||
api_base = "https://generativelanguage.googleapis.com/v1beta"
|
||||
if not api_base.endswith(("/v1", "/v1beta")):
|
||||
api_base = f"{api_base}/v1beta"
|
||||
target_url = f"{api_base}/models/{model}:generateContent"
|
||||
|
||||
image_size = str(context.get_provider_value("image_size", provider)).upper()
|
||||
parts: list[dict[str, Any]] = [{"text": prompt}]
|
||||
for image_bytes, mime in (images_data or [])[:14]:
|
||||
parts.append(
|
||||
{
|
||||
"inlineData": {
|
||||
"mimeType": mime if mime.startswith("image/") else "image/png",
|
||||
"data": base64.b64encode(image_bytes).decode("ascii"),
|
||||
}
|
||||
}
|
||||
)
|
||||
payload = {
|
||||
"contents": [{"role": "user", "parts": parts}],
|
||||
"generationConfig": {
|
||||
"responseModalities": ["TEXT", "IMAGE"],
|
||||
"imageConfig": {
|
||||
"image_size": image_size if image_size in {"1K", "2K", "4K"} else "2K",
|
||||
"aspect_ratio": context.get_provider_value("aspect_ratio", provider),
|
||||
},
|
||||
},
|
||||
}
|
||||
return await context.request_json(
|
||||
target_url,
|
||||
{"x-goog-api-key": context.get_provider_value("api_key", provider)},
|
||||
payload,
|
||||
context.get_provider_value("timeout", provider),
|
||||
"Google Gemini",
|
||||
provider,
|
||||
)
|
||||
@@ -0,0 +1,93 @@
|
||||
import base64
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
|
||||
from ....utils.logger import logger
|
||||
from .context import DrawingRequestContext
|
||||
|
||||
|
||||
async def call_grok_api(
|
||||
context: DrawingRequestContext,
|
||||
prompt: str,
|
||||
images_data: list[tuple[bytes, str]] | None = None,
|
||||
provider: dict | None = None,
|
||||
) -> bytes | None:
|
||||
"""调用 xAI Grok Imagine 官方图片接口。
|
||||
|
||||
Args:
|
||||
prompt: 图片生成或编辑提示词。
|
||||
images_data: 可选参考图片及其 MIME 类型列表,当前只使用第一张。
|
||||
|
||||
Returns:
|
||||
API 返回的图片二进制数据。
|
||||
|
||||
Raises:
|
||||
Exception: 请求失败、响应不是 JSON 或响应中没有有效图片。
|
||||
"""
|
||||
provider = provider or {}
|
||||
raw_url = context.get_provider_value("api_url", provider)
|
||||
target_url = context.build_target_url(raw_url, "grok")
|
||||
api_key = context.get_provider_value("api_key", provider)
|
||||
model = context.get_provider_value("model", provider)
|
||||
timeout = context.get_provider_value("timeout", provider)
|
||||
aspect_ratio = context.get_provider_value("aspect_ratio", provider)
|
||||
|
||||
headers = {
|
||||
"Authorization": f"Bearer {api_key}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
response_format = (
|
||||
context.get_provider_value("response_format", provider) or "b64_json"
|
||||
)
|
||||
payload: dict[str, Any] = {
|
||||
"model": model,
|
||||
"prompt": prompt,
|
||||
"response_format": response_format,
|
||||
}
|
||||
if aspect_ratio:
|
||||
payload["aspect_ratio"] = aspect_ratio
|
||||
|
||||
reference_bytes = 0
|
||||
if images_data:
|
||||
if target_url.endswith("/generations"):
|
||||
target_url = target_url.removesuffix("/generations") + "/edits"
|
||||
image_bytes, mime = images_data[0]
|
||||
image_mime = mime if mime.startswith("image/") else "image/png"
|
||||
encoded = base64.b64encode(image_bytes).decode("ascii")
|
||||
payload["image"] = {
|
||||
"type": "image_url",
|
||||
"url": f"data:{image_mime};base64,{encoded}",
|
||||
}
|
||||
reference_bytes = len(image_bytes)
|
||||
elif target_url.endswith("/edits"):
|
||||
target_url = target_url.removesuffix("/edits") + "/generations"
|
||||
|
||||
logger.info(
|
||||
f"[Comic] 发起 Grok Images API 请求 -> {context.sanitize_url(target_url)} "
|
||||
f"(model={model}, aspect_ratio={aspect_ratio}, "
|
||||
f"reference_bytes={reference_bytes})..."
|
||||
)
|
||||
api_timeout = httpx.Timeout(connect=20.0, read=timeout, write=20.0, pool=20.0)
|
||||
async with httpx.AsyncClient(
|
||||
timeout=api_timeout, proxy=context.get_request_proxy(provider)
|
||||
) as client:
|
||||
resp = await client.post(target_url, headers=headers, json=payload)
|
||||
|
||||
if not 200 <= resp.status_code < 300:
|
||||
error_summary = resp.text[:500] if resp.text else "(空响应)"
|
||||
raise Exception(f"Grok API 请求失败 [HTTP {resp.status_code}]: {error_summary}")
|
||||
|
||||
try:
|
||||
data = resp.json()
|
||||
except Exception:
|
||||
raise Exception(
|
||||
f"Grok API 未返回合法的 JSON [HTTP {resp.status_code}]: "
|
||||
f"<body len={len(resp.content)}>"
|
||||
)
|
||||
|
||||
image = await context.extract_image(data, context.get_request_proxy(provider))
|
||||
if image:
|
||||
return image
|
||||
|
||||
raise Exception(f"Grok API 返回格式异常: {context.summarize_response(data)}")
|
||||
@@ -0,0 +1,187 @@
|
||||
"""OpenAI Images 兼容接口的请求实现。
|
||||
|
||||
该文件负责在同一套 Images API 中区分文生图和图生图请求:没有参考图时使用
|
||||
JSON 的 ``/images/generations``,有参考图时使用 multipart 的 ``/images/edits``。
|
||||
它同时将预设中的 GPT Image 专属参数限制在对应模型和输出格式下,避免兼容端点
|
||||
因未知字段拒绝请求。
|
||||
"""
|
||||
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
|
||||
from ....utils.logger import logger
|
||||
from .context import DrawingRequestContext
|
||||
|
||||
|
||||
async def call_images_api(
|
||||
context: DrawingRequestContext,
|
||||
prompt: str,
|
||||
images_data: list[tuple[bytes, str]] | None = None,
|
||||
provider: dict | None = None,
|
||||
) -> bytes | None:
|
||||
"""调用 OpenAI Images 兼容接口。
|
||||
|
||||
文生图走 JSON,带参考图时切换为 multipart 的 edits 请求。供应商专属
|
||||
参数仅在显式配置时写入,以免把 GPT Image 的参数错误发送给兼容端点。
|
||||
"""
|
||||
provider = provider or {}
|
||||
raw_url = context.get_provider_value("api_url", provider)
|
||||
target_url = context.build_target_url(raw_url, "images")
|
||||
|
||||
api_key = context.get_provider_value("api_key", provider)
|
||||
model = context.get_provider_value("model", provider)
|
||||
timeout = context.get_provider_value("timeout", provider)
|
||||
|
||||
headers = {
|
||||
"Authorization": f"Bearer {api_key}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
|
||||
raw_size = context.get_provider_value("image_size", provider)
|
||||
ar = context.get_provider_value("aspect_ratio", provider)
|
||||
resolved_size = context.resolve_size(raw_size, ar)
|
||||
output_format = context.get_provider_value("output_format", provider)
|
||||
|
||||
payload: dict[str, Any] = {
|
||||
"prompt": prompt,
|
||||
"model": model,
|
||||
"n": 1,
|
||||
"size": resolved_size,
|
||||
"output_format": output_format,
|
||||
}
|
||||
|
||||
# 优先使用预设中的 quality,保留 image_quality 以兼容旧版全局配置。
|
||||
# ``auto`` 只有用户在预设中显式选择时才会发送,避免改变旧版全局配置
|
||||
# 中 ``auto`` 表示“让服务端保持默认”的既有语义。
|
||||
configured_quality = str(provider.get("quality") or "").strip().lower()
|
||||
quality = (
|
||||
configured_quality
|
||||
or str(context.get_provider_value("image_quality", provider) or "")
|
||||
.strip()
|
||||
.lower()
|
||||
)
|
||||
should_send_quality = quality in {"low", "medium", "high"} or (
|
||||
configured_quality == "auto"
|
||||
)
|
||||
if should_send_quality:
|
||||
payload["quality"] = quality
|
||||
|
||||
bg = context.get_provider_value("background", provider)
|
||||
is_gpt_image = str(model).lower().startswith("gpt-image")
|
||||
# background、压缩率和审核策略都是 GPT Image 专属字段。背景为 auto
|
||||
# 时不传,交由服务端按其默认策略处理。
|
||||
if is_gpt_image and bg and bg != "auto":
|
||||
payload["background"] = bg
|
||||
|
||||
# 压缩率仅适用于 JPEG/WebP,审核策略留空时不传,确保其他 OpenAI
|
||||
# 兼容服务不会收到未知字段。
|
||||
response_format = str(provider.get("response_format") or "").strip()
|
||||
if response_format:
|
||||
payload["response_format"] = response_format
|
||||
try:
|
||||
output_compression = int(provider.get("output_compression", 0))
|
||||
except (TypeError, ValueError):
|
||||
output_compression = 0
|
||||
if is_gpt_image and output_compression > 0 and output_format in {"jpeg", "webp"}:
|
||||
payload["output_compression"] = min(output_compression, 100)
|
||||
moderation = str(provider.get("moderation") or "").strip()
|
||||
if is_gpt_image and moderation:
|
||||
payload["moderation"] = moderation
|
||||
|
||||
# 某些中转服务只实现 generations;启用后显式丢弃角色参考图,避免请求
|
||||
# 被自动改写到不支持的 edits 端点。
|
||||
if provider.get("generations_only", False) and images_data:
|
||||
logger.info(
|
||||
"[Comic] OpenAI Images 已启用仅文生图模式,忽略 %d 张参考图。",
|
||||
len(images_data),
|
||||
)
|
||||
images_data = None
|
||||
|
||||
# 模板的参考图数量是用户侧的主动限制。它在切换为 edits 前执行,因此
|
||||
# 设置为 0 会自然退回到文生图;无效输入则使用模板默认值 6。
|
||||
if images_data:
|
||||
try:
|
||||
max_references = max(0, int(provider.get("max_reference_images", 6)))
|
||||
except (TypeError, ValueError):
|
||||
max_references = 6
|
||||
if len(images_data) > max_references:
|
||||
logger.info(
|
||||
"[Comic] OpenAI Images 参考图数量从 %d 张限制为 %d 张。",
|
||||
len(images_data),
|
||||
max_references,
|
||||
)
|
||||
images_data = images_data[:max_references]
|
||||
|
||||
if images_data and len(images_data) > 0:
|
||||
if target_url.endswith("/generations"):
|
||||
target_url = target_url.replace("/generations", "/edits")
|
||||
|
||||
headers.pop(
|
||||
"Content-Type", None
|
||||
) # 移除 JSON 的 Content-Type,让 httpx 自动设置为 multipart/form-data
|
||||
|
||||
multipart_data: dict[str, str] = {
|
||||
"prompt": prompt,
|
||||
"model": model,
|
||||
"n": "1",
|
||||
"size": resolved_size,
|
||||
"output_format": output_format,
|
||||
}
|
||||
if should_send_quality:
|
||||
multipart_data["quality"] = quality
|
||||
if is_gpt_image and bg and bg != "auto":
|
||||
multipart_data["background"] = bg
|
||||
if response_format:
|
||||
multipart_data["response_format"] = response_format
|
||||
if (
|
||||
is_gpt_image
|
||||
and output_compression > 0
|
||||
and output_format in {"jpeg", "webp"}
|
||||
):
|
||||
multipart_data["output_compression"] = str(min(output_compression, 100))
|
||||
if is_gpt_image and moderation:
|
||||
multipart_data["moderation"] = moderation
|
||||
|
||||
files = []
|
||||
for index, (img_bytes, mime) in enumerate(images_data, start=1):
|
||||
ext = mime.split("/")[-1] if "/" in mime else "png"
|
||||
files.append(("image[]", (f"image_{index}.{ext}", img_bytes, mime)))
|
||||
|
||||
logger.info(
|
||||
f"[Comic] 发起 Images API 请求 (含图) -> {context.sanitize_url(target_url)} "
|
||||
f"(model={model}, size={resolved_size}, aspect_ratio={ar}, "
|
||||
f"references={len(images_data)}, reference_bytes={sum(len(image[0]) for image in images_data)})..."
|
||||
)
|
||||
api_timeout = httpx.Timeout(connect=20.0, read=timeout, write=20.0, pool=20.0)
|
||||
async with httpx.AsyncClient(
|
||||
timeout=api_timeout, proxy=context.get_request_proxy(provider)
|
||||
) as client:
|
||||
resp = await client.post(
|
||||
target_url, headers=headers, data=multipart_data, files=files
|
||||
)
|
||||
else:
|
||||
logger.info(
|
||||
f"[Comic] 发起 Images API 请求 -> {context.sanitize_url(target_url)} (model={model}, size={resolved_size}, aspect_ratio={ar})..."
|
||||
)
|
||||
api_timeout = httpx.Timeout(connect=20.0, read=timeout, write=20.0, pool=20.0)
|
||||
async with httpx.AsyncClient(
|
||||
timeout=api_timeout, proxy=context.get_request_proxy(provider)
|
||||
) as client:
|
||||
resp = await client.post(target_url, headers=headers, json=payload)
|
||||
|
||||
if resp.status_code != 200:
|
||||
snippet = resp.text[:500] if resp.text else "(空响应)"
|
||||
raise Exception(f"API 请求失败 [HTTP {resp.status_code}]: {snippet}")
|
||||
|
||||
try:
|
||||
data = resp.json()
|
||||
except Exception:
|
||||
snippet = resp.text[:500] if resp.text else "(空正文)"
|
||||
raise Exception(f"API 未返回合法的 JSON [HTTP {resp.status_code}]: {snippet}")
|
||||
|
||||
image = await context.extract_image(data, context.get_request_proxy(provider))
|
||||
if image:
|
||||
return image
|
||||
|
||||
raise Exception(f"API 返回格式异常: {context.summarize_response(data)}")
|
||||
@@ -0,0 +1,394 @@
|
||||
"""漫画绘图供应商预设的请求体构造。
|
||||
|
||||
本模块只处理各家官方接口不兼容的字段、端点和能力约束;通用 HTTP
|
||||
发送、重试以及图片响应提取仍由上层服务统一负责。这样新增预设时只需
|
||||
在对应分支描述其请求格式,不会把服务商细节重新堆回 DrawingClient。
|
||||
"""
|
||||
|
||||
import base64
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
|
||||
from ....utils.logger import logger
|
||||
from .context import DrawingRequestContext
|
||||
|
||||
|
||||
async def call_preset_api(
|
||||
context: DrawingRequestContext,
|
||||
prompt: str,
|
||||
images_data: list[tuple[bytes, str]] | None,
|
||||
provider: dict,
|
||||
provider_type: str,
|
||||
) -> bytes | None:
|
||||
"""调用使用专有请求格式的绘图供应商预设。
|
||||
|
||||
Args:
|
||||
context: 由绘图客户端提供的公共配置、代理和响应处理能力。
|
||||
prompt: 本次漫画分镜的提示词。
|
||||
images_data: 已读取的角色参考图片二进制数据。
|
||||
provider: 当前供应商模板保存后的配置。
|
||||
provider_type: 预设类型,用于选择官方接口格式。
|
||||
|
||||
Returns:
|
||||
成功时返回图片二进制数据,没有图片时返回 ``None``。
|
||||
|
||||
Raises:
|
||||
ValueError: 供应商预设类型不受支持时抛出。
|
||||
"""
|
||||
api_key = context.get_provider_value("api_key", provider)
|
||||
api_base = str(context.get_provider_value("api_url", provider)).rstrip("/")
|
||||
model = context.get_provider_value("model", provider)
|
||||
timeout = context.get_provider_value("timeout", provider)
|
||||
image_size = str(context.get_provider_value("image_size", provider))
|
||||
aspect_ratio = context.get_provider_value("aspect_ratio", provider)
|
||||
output_format = context.get_provider_value("output_format", provider)
|
||||
# 各个原生 JSON 接口都可接收 data URI。这里统一转换一次,后续分支只
|
||||
# 负责自己的字段语义和参考图数量上限。
|
||||
data_uris = [
|
||||
f"data:{mime if mime.startswith('image/') else 'image/png'};base64,"
|
||||
f"{base64.b64encode(image_bytes).decode('ascii')}"
|
||||
for image_bytes, mime in images_data or []
|
||||
]
|
||||
headers = {"Authorization": f"Bearer {api_key}"}
|
||||
|
||||
if provider_type == "agnes_ai":
|
||||
base = api_base or "https://apihub.agnes-ai.com"
|
||||
target_url = (
|
||||
f"{base}/images/generations"
|
||||
if "/v1" in base
|
||||
else f"{base}/v1/images/generations"
|
||||
)
|
||||
payload: dict[str, Any] = {
|
||||
"model": model,
|
||||
"prompt": prompt,
|
||||
"size": context.resolve_size(image_size, aspect_ratio),
|
||||
"extra_body": {"response_format": output_format or "url"},
|
||||
}
|
||||
if data_uris:
|
||||
payload["extra_body"]["image"] = data_uris
|
||||
provider_name = "Agnes AI"
|
||||
elif provider_type == "xai":
|
||||
base = api_base or "https://api.x.ai"
|
||||
base = base if base.endswith("/v1") else f"{base}/v1"
|
||||
payload = {
|
||||
"model": model,
|
||||
"prompt": prompt,
|
||||
"n": 1,
|
||||
"resolution": image_size.lower()
|
||||
if image_size.upper() in {"1K", "2K"}
|
||||
else "2k",
|
||||
"response_format": output_format or "url",
|
||||
}
|
||||
target_url = f"{base}/images/generations"
|
||||
if data_uris:
|
||||
target_url = f"{base}/images/edits"
|
||||
image_items = [
|
||||
{"type": "image_url", "url": data_uri} for data_uri in data_uris[:5]
|
||||
]
|
||||
payload["image" if len(image_items) == 1 else "images"] = (
|
||||
image_items[0] if len(image_items) == 1 else image_items
|
||||
)
|
||||
payload["aspect_ratio"] = aspect_ratio
|
||||
provider_name = "xAI"
|
||||
elif provider_type == "minimax":
|
||||
base = (api_base or "https://api.minimaxi.com").removesuffix("/v1")
|
||||
target_url = f"{base}/v1/image_generation"
|
||||
payload = {
|
||||
"model": model,
|
||||
"prompt": prompt,
|
||||
"response_format": "url",
|
||||
"n": 1,
|
||||
"aspect_ratio": aspect_ratio,
|
||||
}
|
||||
if data_uris:
|
||||
payload["subject_reference"] = [
|
||||
{"type": "character", "image_file": data_uri}
|
||||
for data_uri in data_uris[:9]
|
||||
]
|
||||
provider_name = "MiniMax"
|
||||
elif provider_type == "doubao":
|
||||
base = api_base or "https://ark.cn-beijing.volces.com"
|
||||
endpoint = (
|
||||
"/api/plan/v3/images/generations"
|
||||
if provider.get("endpoint_mode") == "agent_plan"
|
||||
else "/api/v3/images/generations"
|
||||
)
|
||||
target_url = f"{base}{endpoint}"
|
||||
configured_model = str(provider.get("endpoint_id") or model).strip()
|
||||
model = configured_model or model
|
||||
size_mode = str(provider.get("size_mode") or "preset").strip().lower()
|
||||
configured_size = (
|
||||
provider.get("custom_size")
|
||||
if size_mode == "custom"
|
||||
else provider.get("size") or image_size
|
||||
)
|
||||
resolved_doubao_size = str(configured_size or image_size).strip()
|
||||
if "x" in resolved_doubao_size.lower() or "×" in resolved_doubao_size:
|
||||
resolved_doubao_size = resolved_doubao_size.lower().replace("×", "x")
|
||||
elif resolved_doubao_size.upper() not in {"1K", "2K", "3K", "4K"}:
|
||||
resolved_doubao_size = context.resolve_size(
|
||||
resolved_doubao_size, aspect_ratio
|
||||
)
|
||||
is_seedream_5_pro = (
|
||||
str(provider.get("model_capability") or "").lower() == "seedream_5_pro"
|
||||
or "seedream-5.0-pro" in model.lower()
|
||||
or "seedream-5-0-pro" in model.lower()
|
||||
)
|
||||
payload: dict[str, Any] = {
|
||||
"model": model,
|
||||
"prompt": prompt,
|
||||
"response_format": "url",
|
||||
"output_format": output_format or "png",
|
||||
"watermark": bool(provider.get("watermark", False)),
|
||||
"size": resolved_doubao_size,
|
||||
}
|
||||
if data_uris:
|
||||
max_references = 10 if is_seedream_5_pro else 14
|
||||
try:
|
||||
max_references = min(
|
||||
max_references,
|
||||
max(0, int(provider.get("max_reference_images", max_references))),
|
||||
)
|
||||
except (TypeError, ValueError):
|
||||
pass
|
||||
selected_images = data_uris[:max_references]
|
||||
if selected_images:
|
||||
payload["image"] = (
|
||||
selected_images[0] if len(selected_images) == 1 else selected_images
|
||||
)
|
||||
optimize_mode = str(provider.get("optimize_prompt_mode") or "").strip()
|
||||
if optimize_mode in {"standard", "fast"}:
|
||||
payload["optimize_prompt_options"] = {"mode": optimize_mode}
|
||||
sequential_mode = provider.get("sequential_image_generation")
|
||||
if sequential_mode == "auto" and not is_seedream_5_pro:
|
||||
payload["sequential_image_generation"] = "auto"
|
||||
try:
|
||||
sequential_max_images = int(provider.get("sequential_max_images", 0))
|
||||
except (TypeError, ValueError):
|
||||
sequential_max_images = 0
|
||||
if 1 <= sequential_max_images <= 12:
|
||||
payload["sequential_image_generation_options"] = {
|
||||
"max_images": sequential_max_images
|
||||
}
|
||||
elif sequential_mode == "auto":
|
||||
logger.info("[Comic] 豆包 Seedream 5.0 Pro 不支持组图生成,已忽略该配置。")
|
||||
provider_name = "豆包"
|
||||
elif provider_type == "sensenova":
|
||||
base = api_base or "https://token.sensenova.cn"
|
||||
target_url = (
|
||||
f"{base}/images/generations"
|
||||
if base.endswith("/v1")
|
||||
else f"{base}/v1/images/generations"
|
||||
)
|
||||
if data_uris:
|
||||
logger.info(
|
||||
"[Comic] SenseNova U1 Fast 不支持参考图,已忽略 %d 张。",
|
||||
len(data_uris),
|
||||
)
|
||||
# U1 Fast 仅接受官方枚举尺寸,宽高比不能映射时使用用户配置的默认
|
||||
# 尺寸,仍不合法时再回退到官方横版默认值。
|
||||
size_map = {
|
||||
"1:1": "2048x2048",
|
||||
"2:3": "1664x2496",
|
||||
"3:2": "2496x1664",
|
||||
"16:9": "2752x1536",
|
||||
"9:16": "1536x2752",
|
||||
"4:3": "2368x1760",
|
||||
"3:4": "1760x2368",
|
||||
"4:5": "1824x2272",
|
||||
"5:4": "2272x1824",
|
||||
"21:9": "3072x1376",
|
||||
"9:21": "1344x3136",
|
||||
}
|
||||
default_size = str(provider.get("default_size") or "2752x1536").lower()
|
||||
allowed_sizes = set(size_map.values())
|
||||
resolved_sensenova_size = size_map.get(str(aspect_ratio), default_size)
|
||||
if resolved_sensenova_size not in allowed_sizes:
|
||||
logger.warning(
|
||||
"[Comic] SenseNova 默认尺寸 %s 不受支持,回退为 2752x1536。",
|
||||
resolved_sensenova_size,
|
||||
)
|
||||
resolved_sensenova_size = "2752x1536"
|
||||
try:
|
||||
sensenova_n = max(1, min(4, int(provider.get("n", 1))))
|
||||
except (TypeError, ValueError):
|
||||
sensenova_n = 1
|
||||
payload = {
|
||||
"model": model,
|
||||
"prompt": prompt,
|
||||
"size": resolved_sensenova_size,
|
||||
"n": sensenova_n,
|
||||
}
|
||||
provider_name = "SenseNova"
|
||||
elif provider_type == "dashscope":
|
||||
endpoint_mode = str(provider.get("endpoint_mode", "dashscope"))
|
||||
base = api_base or (
|
||||
"https://token-plan.cn-beijing.maas.aliyuncs.com"
|
||||
if endpoint_mode == "token_plan"
|
||||
else "https://dashscope.aliyuncs.com"
|
||||
)
|
||||
clean_base = base.rstrip("/")
|
||||
if clean_base.endswith("/api/v1"):
|
||||
clean_base = clean_base[:-7].rstrip("/")
|
||||
elif clean_base.endswith("/v1"):
|
||||
clean_base = clean_base[:-3].rstrip("/")
|
||||
target_url = (
|
||||
f"{clean_base}/api/v1/services/aigc/multimodal-generation/generation"
|
||||
)
|
||||
content: list[dict[str, str]] = [{"text": prompt}]
|
||||
try:
|
||||
dashscope_max_references = min(
|
||||
9, max(0, int(provider.get("max_reference_images", 9)))
|
||||
)
|
||||
except (TypeError, ValueError):
|
||||
dashscope_max_references = 9
|
||||
content.extend(
|
||||
{"image": data_uri} for data_uri in data_uris[:dashscope_max_references]
|
||||
)
|
||||
size_mode = str(provider.get("size_mode") or "preset").strip().lower()
|
||||
custom_size = str(provider.get("custom_size") or "").strip()
|
||||
if size_mode == "custom" and custom_size:
|
||||
normalized_size = custom_size.upper()
|
||||
if normalized_size not in {"1K", "2K", "4K"}:
|
||||
normalized_size = (
|
||||
custom_size.lower().replace("×", "*").replace("x", "*")
|
||||
)
|
||||
dashscope_size = normalized_size
|
||||
else:
|
||||
dashscope_size = resolve_dashscope_size(image_size, aspect_ratio)
|
||||
is_wan27 = str(model).startswith("wan2.7")
|
||||
enable_sequential = is_wan27 and bool(provider.get("enable_sequential", False))
|
||||
if enable_sequential:
|
||||
dashscope_n_limit = 12
|
||||
elif is_wan27:
|
||||
dashscope_n_limit = 4
|
||||
elif str(model).startswith("qwen-image-2.0"):
|
||||
dashscope_n_limit = 6
|
||||
else:
|
||||
dashscope_n_limit = 1
|
||||
try:
|
||||
dashscope_n = max(1, min(dashscope_n_limit, int(provider.get("n", 1))))
|
||||
except (TypeError, ValueError):
|
||||
dashscope_n = 1
|
||||
parameters: dict[str, Any] = {
|
||||
"size": dashscope_size,
|
||||
"n": dashscope_n,
|
||||
"watermark": bool(provider.get("watermark", False)),
|
||||
}
|
||||
negative_prompt = str(provider.get("negative_prompt") or "").strip()
|
||||
if negative_prompt and not is_wan27:
|
||||
parameters["negative_prompt"] = negative_prompt
|
||||
elif negative_prompt:
|
||||
logger.info("[Comic] DashScope wan2.7 不支持负面提示词,已忽略该配置。")
|
||||
if is_wan27:
|
||||
if enable_sequential:
|
||||
parameters["enable_sequential"] = True
|
||||
else:
|
||||
parameters["thinking_mode"] = bool(provider.get("thinking_mode", True))
|
||||
else:
|
||||
parameters["prompt_extend"] = bool(provider.get("prompt_extend", False))
|
||||
payload = {
|
||||
"model": model,
|
||||
"input": {"messages": [{"role": "user", "content": content}]},
|
||||
"parameters": parameters,
|
||||
}
|
||||
provider_name = "DashScope"
|
||||
elif provider_type == "stepfun":
|
||||
return await call_stepfun_api(
|
||||
context, prompt, images_data, provider, api_key, model, timeout
|
||||
)
|
||||
else:
|
||||
raise ValueError(f"不支持的绘图供应商预设: {provider_type}")
|
||||
|
||||
return await context.request_json(
|
||||
target_url, headers, payload, timeout, provider_name, provider
|
||||
)
|
||||
|
||||
|
||||
async def call_stepfun_api(
|
||||
context: DrawingRequestContext,
|
||||
prompt: str,
|
||||
images_data: list[tuple[bytes, str]] | None,
|
||||
provider: dict,
|
||||
api_key: str,
|
||||
model: str,
|
||||
timeout: int | float,
|
||||
) -> bytes | None:
|
||||
"""调用阶跃星辰图片接口,图生图使用官方 multipart 字段。"""
|
||||
target_url = context.build_target_url(
|
||||
context.get_provider_value("api_url", provider), "images"
|
||||
)
|
||||
headers = {"Authorization": f"Bearer {api_key}"}
|
||||
api_timeout = httpx.Timeout(connect=20.0, read=timeout, write=20.0, pool=20.0)
|
||||
|
||||
if images_data:
|
||||
target_url = target_url.replace("/generations", "/edits")
|
||||
image_bytes, mime = images_data[0]
|
||||
extension = mime.split("/")[-1] if "/" in mime else "png"
|
||||
form_data = {
|
||||
"model": model,
|
||||
"prompt": prompt,
|
||||
"response_format": "url",
|
||||
}
|
||||
files = {
|
||||
"image": (f"reference.{extension}", image_bytes, mime),
|
||||
}
|
||||
logger.info(
|
||||
f"[Comic] 发起阶跃星辰图生图请求 -> {context.sanitize_url(target_url)}"
|
||||
)
|
||||
async with httpx.AsyncClient(
|
||||
timeout=api_timeout, proxy=context.get_request_proxy(provider)
|
||||
) as client:
|
||||
response = await client.post(
|
||||
target_url, headers=headers, data=form_data, files=files
|
||||
)
|
||||
else:
|
||||
payload = {
|
||||
"model": model,
|
||||
"prompt": prompt,
|
||||
"size": context.resolve_size(
|
||||
context.get_provider_value("image_size", provider),
|
||||
context.get_provider_value("aspect_ratio", provider),
|
||||
),
|
||||
"response_format": "url",
|
||||
}
|
||||
logger.info(
|
||||
f"[Comic] 发起阶跃星辰文生图请求 -> {context.sanitize_url(target_url)}"
|
||||
)
|
||||
async with httpx.AsyncClient(
|
||||
timeout=api_timeout, proxy=context.get_request_proxy(provider)
|
||||
) as client:
|
||||
response = await client.post(target_url, headers=headers, json=payload)
|
||||
|
||||
if not 200 <= response.status_code < 300:
|
||||
message = response.text[:500] if response.text else "(空响应)"
|
||||
raise Exception(
|
||||
f"阶跃星辰 API 请求失败 [HTTP {response.status_code}]: {message}"
|
||||
)
|
||||
try:
|
||||
data = response.json()
|
||||
except ValueError as exc:
|
||||
raise Exception("阶跃星辰 API 未返回合法 JSON") from exc
|
||||
image = await context.extract_image(data, context.get_request_proxy(provider))
|
||||
if image:
|
||||
return image
|
||||
raise Exception(f"阶跃星辰 API 返回格式异常: {context.summarize_response(data)}")
|
||||
|
||||
|
||||
def resolve_dashscope_size(image_size: str, aspect_ratio: str) -> str:
|
||||
"""将漫画尺寸和比例换算为 DashScope 的 size 格式。"""
|
||||
long_edge = {"1K": 1280, "2K": 2048, "4K": 4096}.get(image_size.upper(), 2048)
|
||||
try:
|
||||
width_ratio, height_ratio = (int(value) for value in aspect_ratio.split(":", 1))
|
||||
if width_ratio <= 0 or height_ratio <= 0:
|
||||
raise ValueError
|
||||
except (AttributeError, TypeError, ValueError):
|
||||
width_ratio, height_ratio = 1, 1
|
||||
if width_ratio >= height_ratio:
|
||||
width = long_edge
|
||||
height = round(long_edge * height_ratio / width_ratio / 16) * 16
|
||||
else:
|
||||
height = long_edge
|
||||
width = round(long_edge * width_ratio / height_ratio / 16) * 16
|
||||
return f"{max(512, width)}*{max(512, height)}"
|
||||
@@ -0,0 +1,114 @@
|
||||
"""绘图供应商请求服务的统一分发层。
|
||||
|
||||
每个服务商协议位于独立模块中,本服务仅提供稳定的调用入口并传递共享上下文。
|
||||
``DrawingClient`` 因此不需要了解某个服务商的请求体格式,也无需使用 Mixin
|
||||
将大量实现混入客户端类。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
from .chat import call_chat_api
|
||||
from .common import post_json_for_image
|
||||
from .context import DrawingRequestContext
|
||||
from .gemini import call_gemini_api
|
||||
from .google import call_google_api
|
||||
from .grok import call_grok_api
|
||||
from .images import call_images_api
|
||||
from .presets import call_preset_api, call_stepfun_api
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class DrawingApiRequestService:
|
||||
"""协调各服务商请求实现并对外暴露统一调用入口。
|
||||
|
||||
每个方法保持一层直接转发,目的是固定高层调用契约;服务商特有的参数和
|
||||
能力限制仍留在各自模块,避免该类成为新的请求逻辑聚集点。
|
||||
"""
|
||||
|
||||
context: DrawingRequestContext
|
||||
|
||||
async def call_google_api(
|
||||
self,
|
||||
prompt: str,
|
||||
images_data: list[tuple[bytes, str]] | None,
|
||||
provider: dict,
|
||||
) -> bytes | None:
|
||||
return await call_google_api(self.context, prompt, images_data, provider)
|
||||
|
||||
async def call_preset_api(
|
||||
self,
|
||||
prompt: str,
|
||||
images_data: list[tuple[bytes, str]] | None,
|
||||
provider: dict,
|
||||
provider_type: str,
|
||||
) -> bytes | None:
|
||||
return await call_preset_api(
|
||||
self.context, prompt, images_data, provider, provider_type
|
||||
)
|
||||
|
||||
async def post_json_for_image(
|
||||
self,
|
||||
target_url: str,
|
||||
headers: dict[str, str],
|
||||
payload: dict[str, Any],
|
||||
timeout: int | float,
|
||||
provider_name: str,
|
||||
provider: dict,
|
||||
) -> bytes | None:
|
||||
return await post_json_for_image(
|
||||
self.context,
|
||||
target_url,
|
||||
headers,
|
||||
payload,
|
||||
timeout,
|
||||
provider_name,
|
||||
provider,
|
||||
)
|
||||
|
||||
async def call_stepfun_api(
|
||||
self,
|
||||
prompt: str,
|
||||
images_data: list[tuple[bytes, str]] | None,
|
||||
provider: dict,
|
||||
api_key: str,
|
||||
model: str,
|
||||
timeout: int | float,
|
||||
) -> bytes | None:
|
||||
return await call_stepfun_api(
|
||||
self.context, prompt, images_data, provider, api_key, model, timeout
|
||||
)
|
||||
|
||||
async def call_images_api(
|
||||
self,
|
||||
prompt: str,
|
||||
images_data: list[tuple[bytes, str]] | None = None,
|
||||
provider: dict | None = None,
|
||||
) -> bytes | None:
|
||||
return await call_images_api(self.context, prompt, images_data, provider)
|
||||
|
||||
async def call_grok_api(
|
||||
self,
|
||||
prompt: str,
|
||||
images_data: list[tuple[bytes, str]] | None = None,
|
||||
provider: dict | None = None,
|
||||
) -> bytes | None:
|
||||
return await call_grok_api(self.context, prompt, images_data, provider)
|
||||
|
||||
async def call_gemini_api(
|
||||
self,
|
||||
prompt: str,
|
||||
images_data: list[tuple[bytes, str]] | None = None,
|
||||
provider: dict | None = None,
|
||||
) -> bytes | None:
|
||||
return await call_gemini_api(self.context, prompt, images_data, provider)
|
||||
|
||||
async def call_chat_api(
|
||||
self,
|
||||
prompt: str,
|
||||
images_data: list[tuple[bytes, str]] | None = None,
|
||||
provider: dict | None = None,
|
||||
) -> bytes | None:
|
||||
return await call_chat_api(self.context, prompt, images_data, provider)
|
||||
@@ -0,0 +1,432 @@
|
||||
"""漫画绘图客户端的高层调度入口。
|
||||
|
||||
``DrawingClient`` 保留插件既有的私有兼容入口,供调用方和扩展继续使用;实际
|
||||
供应商请求与图片响应分别委托给两个组合服务。客户端只保留供应商回退、重试和
|
||||
全局配置读取,避免 HTTP 协议细节再次集中到一个过大的类中。
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import re
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
|
||||
from ...utils.logger import logger
|
||||
from ..config.config_manager import ConfigManager
|
||||
from .api_requests import DrawingApiRequestService
|
||||
from .api_requests.context import DrawingRequestContext
|
||||
from .api_requests.presets import resolve_dashscope_size
|
||||
from .drawing_image_response import (
|
||||
DrawingImageResponseService,
|
||||
ImageDownloadFailedError,
|
||||
)
|
||||
|
||||
__all__ = ["DrawingClient", "ImageDownloadFailedError"]
|
||||
|
||||
|
||||
# 供应商条目是唯一的连接配置来源。这里的默认值只用于兼容手工编辑的缺失字段,
|
||||
# 不再读取配置面板中已移除的外层绘图参数。
|
||||
DRAWING_PROVIDER_DEFAULTS: dict[str, Any] = {
|
||||
"api_url": "",
|
||||
"api_key": "",
|
||||
"model": "gpt-image-2",
|
||||
"api_protocol": "images",
|
||||
"image_size": "1024x1024",
|
||||
"aspect_ratio": "16:9",
|
||||
"image_quality": "high",
|
||||
"background": "auto",
|
||||
"output_format": "png",
|
||||
"timeout": 600,
|
||||
}
|
||||
|
||||
|
||||
class DrawingClient:
|
||||
"""协调绘图供应商选择、重试、请求服务和图片响应处理。
|
||||
|
||||
服务对象通过 ``DrawingRequestContext`` 显式获取所需能力,而不是继承
|
||||
客户端内部状态。私有兼容入口仍保留为薄转发层,确保既有扩展和测试替换
|
||||
``_post_json_for_image`` 等方法时能够继续生效。
|
||||
"""
|
||||
|
||||
def __init__(self, config_manager: ConfigManager):
|
||||
self.config_manager = config_manager
|
||||
self._image_response_service = DrawingImageResponseService(
|
||||
hooks=self,
|
||||
# 保持实例替换下载方法时,响应服务也会使用替换后的实现。
|
||||
download_image=lambda url, proxy: self.download_public_image(url, proxy),
|
||||
)
|
||||
self._request_service = DrawingApiRequestService(
|
||||
DrawingRequestContext(
|
||||
hooks=self,
|
||||
# 保持既有测试和扩展对兼容入口的动态替换能力。
|
||||
request_json=lambda *args: self._post_json_for_image(*args),
|
||||
extract_image=lambda data, proxy: self._extract_image_from_response(
|
||||
data, proxy
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
def _build_target_url(self, raw_url: str, protocol: str) -> str:
|
||||
"""智能解析补全用户配置的 API URL。"""
|
||||
url = (raw_url or "").strip().rstrip("/")
|
||||
if not url:
|
||||
if protocol == "grok":
|
||||
url = "https://api.x.ai"
|
||||
elif protocol == "gemini":
|
||||
url = "https://generativelanguage.googleapis.com"
|
||||
else:
|
||||
url = "https://api.openai.com/v1"
|
||||
|
||||
if protocol == "images":
|
||||
if url.endswith("/images/generations"):
|
||||
return url
|
||||
if url.endswith("/v1"):
|
||||
return f"{url}/images/generations"
|
||||
if "/v1/" in url:
|
||||
return url if "images" in url else f"{url}/images/generations"
|
||||
return f"{url}/v1/images/generations"
|
||||
if protocol == "chat":
|
||||
if url.endswith("/chat/completions"):
|
||||
return url
|
||||
if url.endswith("/v1"):
|
||||
return f"{url}/chat/completions"
|
||||
if "/v1/" in url:
|
||||
return url if "chat" in url else f"{url}/chat/completions"
|
||||
return f"{url}/v1/chat/completions"
|
||||
if protocol == "grok":
|
||||
if url.endswith(("/images/generations", "/images/edits")):
|
||||
return url
|
||||
if url.endswith("/v1"):
|
||||
return f"{url}/images/generations"
|
||||
if "/v1/" in url:
|
||||
return url if "/images/" in url else f"{url}/images/generations"
|
||||
return f"{url}/v1/images/generations"
|
||||
if protocol == "gemini":
|
||||
if url.endswith("/interactions"):
|
||||
return url
|
||||
if url.endswith(("/v1beta", "/v1")):
|
||||
return f"{url}/interactions"
|
||||
return f"{url}/v1beta/interactions"
|
||||
raise ValueError(f"不支持的绘图 API 协议: {protocol}")
|
||||
|
||||
async def generate_image(
|
||||
self,
|
||||
prompt: str,
|
||||
images_data: list[tuple[bytes, str]] | None = None,
|
||||
disable_retry: bool = False,
|
||||
) -> tuple[bytes | None, str | None]:
|
||||
"""调用候选供应商生成单张图片。
|
||||
|
||||
Args:
|
||||
prompt: 用于生成图片的提示词。
|
||||
images_data: 可选参考图片及其 MIME 类型。
|
||||
disable_retry: 是否禁用请求失败后的重试。
|
||||
|
||||
Returns:
|
||||
图片二进制数据与最后一次错误信息组成的元组。
|
||||
"""
|
||||
provider_configs = self.config_manager.get_drawing_provider_configs()
|
||||
if not provider_configs:
|
||||
message = "未配置有效的漫画绘图供应商,请在绘图供应商配置表中添加条目。"
|
||||
logger.warning("[Comic] %s", message)
|
||||
return None, message
|
||||
|
||||
last_error_msg = None
|
||||
last_download_error: ImageDownloadFailedError | None = None
|
||||
for provider in provider_configs:
|
||||
try:
|
||||
result, last_error_msg = await self._generate_image_with_provider(
|
||||
prompt, images_data, disable_retry, provider
|
||||
)
|
||||
except ImageDownloadFailedError as exc:
|
||||
# 图片已经由上游生成,但当前候选返回的 URL 无法下载;继续尝试后备供应商。
|
||||
last_download_error = exc
|
||||
last_error_msg = str(exc)
|
||||
result = None
|
||||
if result:
|
||||
return result, None
|
||||
provider_name = str(provider.get("name", "unnamed")).strip()
|
||||
logger.warning(
|
||||
"[Comic] 绘图供应商 %s 失败,尝试下一个候选。",
|
||||
provider_name or "unnamed",
|
||||
)
|
||||
if last_download_error:
|
||||
raise last_download_error
|
||||
return None, last_error_msg
|
||||
|
||||
async def _generate_image_with_provider(
|
||||
self,
|
||||
prompt: str,
|
||||
images_data: list[tuple[bytes, str]] | None,
|
||||
disable_retry: bool,
|
||||
provider: dict,
|
||||
) -> tuple[bytes | None, str | None]:
|
||||
"""使用一个已配置供应商执行生成并处理重试。"""
|
||||
api_protocol = self._get_provider_value("api_protocol", provider)
|
||||
max_retries = self.config_manager.get_drawing_network_retries()
|
||||
output_exception_retries = (
|
||||
0
|
||||
if disable_retry
|
||||
else self.config_manager.get_drawing_output_exception_retries()
|
||||
)
|
||||
exception_keywords = (
|
||||
self.config_manager.get_drawing_output_exception_retry_keywords()
|
||||
)
|
||||
retry_delay = self.config_manager.get_drawing_retry_delay()
|
||||
exception_retry_count = 0
|
||||
network_retry_count = 0
|
||||
last_error_msg = None
|
||||
|
||||
while True:
|
||||
try:
|
||||
if api_protocol == "images":
|
||||
result = await self._call_images_api(prompt, images_data, provider)
|
||||
elif api_protocol == "chat":
|
||||
result = await self._call_chat_api(prompt, images_data, provider)
|
||||
elif api_protocol == "google":
|
||||
result = await self._call_google_api(prompt, images_data, provider)
|
||||
elif api_protocol == "grok":
|
||||
result = await self._call_grok_api(prompt, images_data, provider)
|
||||
elif api_protocol == "gemini":
|
||||
result = await self._call_gemini_api(prompt, images_data, provider)
|
||||
elif api_protocol in {
|
||||
"agnes_ai",
|
||||
"xai",
|
||||
"minimax",
|
||||
"doubao",
|
||||
"sensenova",
|
||||
"dashscope",
|
||||
"stepfun",
|
||||
}:
|
||||
result = await self._call_preset_api(
|
||||
prompt, images_data, provider, api_protocol
|
||||
)
|
||||
else:
|
||||
raise ValueError(f"不支持的绘图 API 协议: {api_protocol}")
|
||||
if result:
|
||||
return result, None
|
||||
break
|
||||
except Exception as exc:
|
||||
if isinstance(exc, ImageDownloadFailedError):
|
||||
raise
|
||||
last_error_msg = str(exc)
|
||||
logger.error(
|
||||
"[Comic] 画图报错 (%s): %s", type(exc).__name__, last_error_msg
|
||||
)
|
||||
if disable_retry:
|
||||
break
|
||||
is_exception = any(
|
||||
keyword in last_error_msg
|
||||
for keyword in exception_keywords
|
||||
if keyword
|
||||
)
|
||||
if is_exception:
|
||||
if exception_retry_count < output_exception_retries:
|
||||
exception_retry_count += 1
|
||||
logger.info(
|
||||
"[Comic] 命中异常关键词,开始第 %d 次内容重试...",
|
||||
exception_retry_count,
|
||||
)
|
||||
await asyncio.sleep(retry_delay)
|
||||
continue
|
||||
break
|
||||
status_match = re.search(r"HTTP (\d{3})", last_error_msg)
|
||||
status_code = int(status_match.group(1)) if status_match else None
|
||||
is_retryable_network_error = isinstance(exc, httpx.RequestError) or (
|
||||
status_code in {408, 409, 429}
|
||||
or status_code is not None
|
||||
and status_code >= 500
|
||||
)
|
||||
if not is_retryable_network_error or network_retry_count >= max_retries:
|
||||
break
|
||||
network_retry_count += 1
|
||||
logger.info(
|
||||
"[Comic] 网络或服务报错,开始第 %d 次网络重试...",
|
||||
network_retry_count,
|
||||
)
|
||||
await asyncio.sleep(retry_delay)
|
||||
|
||||
logger.debug("[Comic] 画图重试次数耗尽或请求失败,任务终止。")
|
||||
return None, last_error_msg
|
||||
|
||||
def _get_provider_value(self, name: str, provider: dict) -> Any:
|
||||
"""读取供应商条目字段,并为缺失字段提供安全默认值。
|
||||
|
||||
绘图连接参数只允许来自当前条目,避免删除面板外层字段后仍出现不可见的
|
||||
回退来源。模板正常保存时会提供完整字段,默认值仅保护旧数据或手工配置。
|
||||
"""
|
||||
value = provider.get(name)
|
||||
if value not in (None, ""):
|
||||
return value
|
||||
return DRAWING_PROVIDER_DEFAULTS.get(name, "")
|
||||
|
||||
def _get_request_proxy(self, provider: dict | None = None) -> str | None:
|
||||
"""获取当前绘图请求的代理,供应商配置优先于全局配置。"""
|
||||
provider_proxy = str((provider or {}).get("proxy", "")).strip()
|
||||
if provider_proxy:
|
||||
return provider_proxy
|
||||
getter = getattr(self.config_manager, "get_drawing_proxy", None)
|
||||
global_proxy = getter() if callable(getter) else ""
|
||||
return str(global_proxy).strip() or None
|
||||
|
||||
def _resolve_size(self, size_or_ratio: str, aspect_ratio: str) -> str:
|
||||
"""按当前供应商条目的比例解析 API 支持的 WxH 尺寸。"""
|
||||
size = (size_or_ratio or "").strip().lower()
|
||||
aspect_ratio = (aspect_ratio or "").strip().lower()
|
||||
if not aspect_ratio:
|
||||
aspect_ratio = "16:9"
|
||||
size_aliases = {"1k": 1024, "2k": 2560, "4k": 3840}
|
||||
if size in size_aliases:
|
||||
result = self._build_size_from_ratio(size_aliases[size], aspect_ratio)
|
||||
elif size in {"auto", ""}:
|
||||
result = self._build_size_from_ratio(1792, aspect_ratio)
|
||||
elif ":" in size and re.fullmatch(r"\d+:\d+", size):
|
||||
result = self._build_size_from_ratio(1792, size)
|
||||
elif re.fullmatch(r"\d+x\d+", size):
|
||||
result = size
|
||||
else:
|
||||
result = self._build_size_from_ratio(1792, aspect_ratio)
|
||||
if re.fullmatch(r"\d+x\d+", result):
|
||||
try:
|
||||
width, height = map(int, result.split("x"))
|
||||
result = (
|
||||
f"{max(16, ((width + 15) // 16) * 16)}"
|
||||
f"x{max(16, ((height + 15) // 16) * 16)}"
|
||||
)
|
||||
except ValueError:
|
||||
pass
|
||||
return result
|
||||
|
||||
@staticmethod
|
||||
def _build_size_from_ratio(long_edge: int, aspect_ratio: str) -> str:
|
||||
"""按长边和宽高比构建 16 的倍数尺寸。"""
|
||||
if not aspect_ratio or ":" not in aspect_ratio:
|
||||
aspect_ratio = "16:9"
|
||||
try:
|
||||
width_ratio, height_ratio = map(int, aspect_ratio.split(":", 1))
|
||||
except ValueError:
|
||||
width_ratio, height_ratio = 16, 9
|
||||
if width_ratio >= height_ratio:
|
||||
width = long_edge
|
||||
height = max(2, round(long_edge * height_ratio / width_ratio))
|
||||
else:
|
||||
height = long_edge
|
||||
width = max(2, round(long_edge * width_ratio / height_ratio))
|
||||
width = max(16, ((width + 15) // 16) * 16)
|
||||
height = max(16, ((height + 15) // 16) * 16)
|
||||
return f"{width}x{height}"
|
||||
|
||||
async def _call_google_api(
|
||||
self, prompt: str, images_data: list[tuple[bytes, str]] | None, provider: dict
|
||||
) -> bytes | None:
|
||||
return await self._request_service.call_google_api(
|
||||
prompt, images_data, provider
|
||||
)
|
||||
|
||||
async def _call_preset_api(
|
||||
self,
|
||||
prompt: str,
|
||||
images_data: list[tuple[bytes, str]] | None,
|
||||
provider: dict,
|
||||
provider_type: str,
|
||||
) -> bytes | None:
|
||||
return await self._request_service.call_preset_api(
|
||||
prompt, images_data, provider, provider_type
|
||||
)
|
||||
|
||||
async def _call_stepfun_api(
|
||||
self,
|
||||
prompt: str,
|
||||
images_data: list[tuple[bytes, str]] | None,
|
||||
provider: dict,
|
||||
api_key: str,
|
||||
model: str,
|
||||
timeout: int | float,
|
||||
) -> bytes | None:
|
||||
return await self._request_service.call_stepfun_api(
|
||||
prompt, images_data, provider, api_key, model, timeout
|
||||
)
|
||||
|
||||
async def _post_json_for_image(
|
||||
self,
|
||||
target_url: str,
|
||||
headers: dict[str, str],
|
||||
payload: dict[str, Any],
|
||||
timeout: int | float,
|
||||
provider_name: str,
|
||||
provider: dict,
|
||||
) -> bytes | None:
|
||||
return await self._request_service.post_json_for_image(
|
||||
target_url, headers, payload, timeout, provider_name, provider
|
||||
)
|
||||
|
||||
def _resolve_dashscope_size(self, image_size: str, aspect_ratio: str) -> str:
|
||||
return resolve_dashscope_size(image_size, aspect_ratio)
|
||||
|
||||
async def _call_images_api(
|
||||
self,
|
||||
prompt: str,
|
||||
images_data: list[tuple[bytes, str]] | None = None,
|
||||
provider: dict | None = None,
|
||||
) -> bytes | None:
|
||||
return await self._request_service.call_images_api(
|
||||
prompt, images_data, provider
|
||||
)
|
||||
|
||||
async def _call_grok_api(
|
||||
self,
|
||||
prompt: str,
|
||||
images_data: list[tuple[bytes, str]] | None = None,
|
||||
provider: dict | None = None,
|
||||
) -> bytes | None:
|
||||
return await self._request_service.call_grok_api(prompt, images_data, provider)
|
||||
|
||||
async def _call_gemini_api(
|
||||
self,
|
||||
prompt: str,
|
||||
images_data: list[tuple[bytes, str]] | None = None,
|
||||
provider: dict | None = None,
|
||||
) -> bytes | None:
|
||||
return await self._request_service.call_gemini_api(
|
||||
prompt, images_data, provider
|
||||
)
|
||||
|
||||
async def _call_chat_api(
|
||||
self,
|
||||
prompt: str,
|
||||
images_data: list[tuple[bytes, str]] | None = None,
|
||||
provider: dict | None = None,
|
||||
) -> bytes | None:
|
||||
return await self._request_service.call_chat_api(prompt, images_data, provider)
|
||||
|
||||
async def _extract_image_from_response(
|
||||
self, data: Any, proxy: str | None = None
|
||||
) -> bytes | None:
|
||||
return await self._image_response_service.extract_image_from_response(
|
||||
data, proxy
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _decode_data_uri(data_uri: str) -> bytes:
|
||||
return DrawingImageResponseService.decode_data_uri(data_uri)
|
||||
|
||||
@staticmethod
|
||||
def _decode_base64(encoded: str) -> bytes:
|
||||
return DrawingImageResponseService.decode_base64(encoded)
|
||||
|
||||
@staticmethod
|
||||
def _validate_image_bytes(data: bytes) -> None:
|
||||
DrawingImageResponseService.validate_image_bytes(data)
|
||||
|
||||
async def download_public_image(
|
||||
self, url: str, proxy: str | None = None
|
||||
) -> bytes | None:
|
||||
return await self._image_response_service.download_public_image(url, proxy)
|
||||
|
||||
@staticmethod
|
||||
def _sanitize_url(url: str) -> str:
|
||||
return DrawingImageResponseService.sanitize_url(url)
|
||||
|
||||
@staticmethod
|
||||
def _summarize_response(data: Any) -> str:
|
||||
return DrawingImageResponseService.summarize_response(data)
|
||||
@@ -0,0 +1,305 @@
|
||||
"""绘图接口响应的图片提取、下载与安全校验。
|
||||
|
||||
不同供应商会把图片放在 Base64、Data URI、Markdown 或 URL 字段中。本服务按
|
||||
可靠性由高到低依次处理这些候选内容,并对公网下载地址及图片签名做校验,防止
|
||||
将接口错误页、内网地址或超大内容作为漫画图片继续投递。
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import base64
|
||||
import binascii
|
||||
import re
|
||||
from collections.abc import Awaitable, Callable
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Protocol
|
||||
from urllib.parse import urlsplit, urlunsplit
|
||||
|
||||
import httpx
|
||||
|
||||
from ...utils.logger import logger
|
||||
|
||||
|
||||
class ImageDownloadFailedError(Exception):
|
||||
"""图片下载失败,但保留了最后一次尝试的原始 URL 供兜底发送。"""
|
||||
|
||||
def __init__(self, message: str, fallback_url: str | None = None):
|
||||
super().__init__(message)
|
||||
self.fallback_url = fallback_url
|
||||
|
||||
|
||||
class DrawingImageResponseHooks(Protocol):
|
||||
"""描述图片响应服务依赖的宿主能力。"""
|
||||
|
||||
config_manager: Any
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class DrawingImageResponseService:
|
||||
"""封装绘图响应解析、图片下载和安全校验。
|
||||
|
||||
``hooks`` 只保存下载代理所需的配置访问能力,实际下载函数显式注入,使
|
||||
响应处理独立于 ``DrawingClient`` 的具体实现,也允许测试安全地替换下载。
|
||||
"""
|
||||
|
||||
hooks: DrawingImageResponseHooks
|
||||
download_image: Callable[[str, str | None], Awaitable[bytes | None]]
|
||||
|
||||
MAX_IMAGE_BYTES = 100 * 1024 * 1024
|
||||
MAX_IMAGE_REDIRECTS = 5
|
||||
IMAGE_DOWNLOAD_TOTAL_TIMEOUT = 90
|
||||
|
||||
async def extract_image_from_response(
|
||||
self, data: Any, proxy: str | None = None
|
||||
) -> bytes | None:
|
||||
"""递归提取绘图响应中的图片数据。"""
|
||||
encoded: list[tuple[str, str]] = []
|
||||
image_fields: list[tuple[str, str]] = []
|
||||
content_images: list[tuple[str, str]] = []
|
||||
content_urls: list[tuple[str, str]] = []
|
||||
fallback_urls: list[tuple[str, str]] = []
|
||||
|
||||
def collect(value: Any, path: tuple[str, ...] = ()) -> None:
|
||||
if isinstance(value, dict):
|
||||
for name, item in value.items():
|
||||
collect(item, (*path, name.lower()))
|
||||
elif isinstance(value, list):
|
||||
for item in value:
|
||||
collect(item, path)
|
||||
elif isinstance(value, str):
|
||||
text = value.strip()
|
||||
if not text:
|
||||
return
|
||||
key = path[-1] if path else ""
|
||||
if key in {"b64_json", "base64"}:
|
||||
encoded.append(("base64", text))
|
||||
return
|
||||
if key in {"image_url", "image"} or (
|
||||
key == "url"
|
||||
and any(
|
||||
name in {"data", "image", "images", "image_url"}
|
||||
for name in path[:-1]
|
||||
)
|
||||
):
|
||||
image_fields.append(("value", text))
|
||||
return
|
||||
|
||||
data_uris = re.findall(
|
||||
r"data:image/[^\s,;]+(?:;[^\s,;]+)*;base64,[A-Za-z0-9+/=_-]+",
|
||||
text,
|
||||
re.IGNORECASE,
|
||||
)
|
||||
content_images.extend(("value", item) for item in data_uris)
|
||||
|
||||
markdown_urls = re.findall(
|
||||
r"!\[[^\]]*\]\((https?://[^\s<>\"')\]]+)\)", text
|
||||
)
|
||||
content_images.extend(
|
||||
("url", item.rstrip(".,;`")) for item in markdown_urls
|
||||
)
|
||||
|
||||
urls = re.findall(r"https?://[^\s<>\"')\]]+", text)
|
||||
markdown_url_set = set(markdown_urls)
|
||||
target = content_urls if key in {"content", "text"} else fallback_urls
|
||||
target.extend(
|
||||
("url", item.rstrip(".,;`"))
|
||||
for item in urls
|
||||
if item not in markdown_url_set
|
||||
)
|
||||
|
||||
if not data_uris and not urls and len(text) >= 100:
|
||||
encoded.append(("base64", text))
|
||||
|
||||
collect(data)
|
||||
last_download_error: Exception | None = None
|
||||
last_download_url: str | None = None
|
||||
candidates = (
|
||||
encoded + image_fields + content_images + content_urls + fallback_urls
|
||||
)
|
||||
for candidate_type, candidate in candidates:
|
||||
try:
|
||||
if candidate_type == "url" or candidate.startswith(
|
||||
("http://", "https://")
|
||||
):
|
||||
last_download_url = candidate
|
||||
image = await self.download_image(candidate, proxy)
|
||||
elif candidate.startswith("data:image/"):
|
||||
image = self.decode_data_uri(candidate)
|
||||
elif candidate.startswith("base64://"):
|
||||
image = self.decode_base64(candidate[len("base64://") :])
|
||||
else:
|
||||
image = self.decode_base64(candidate)
|
||||
if image:
|
||||
return image
|
||||
except (httpx.HTTPError, httpx.TimeoutException) as exc:
|
||||
logger.warning("[Comic] 图片下载失败 (%s): %s", type(exc).__name__, exc)
|
||||
last_download_error = exc
|
||||
except (ValueError, TypeError, binascii.Error) as exc:
|
||||
logger.warning("[Comic] 跳过无效图片候选内容: %s", exc)
|
||||
|
||||
if last_download_error:
|
||||
raise ImageDownloadFailedError(
|
||||
str(last_download_error), fallback_url=last_download_url
|
||||
)
|
||||
return None
|
||||
|
||||
@classmethod
|
||||
def decode_data_uri(cls, data_uri: str) -> bytes:
|
||||
"""解码 image/* Data URI。"""
|
||||
header, encoded = data_uri.split(",", 1)
|
||||
if ";base64" not in header.lower():
|
||||
raise ValueError("Data URI 不是 Base64 图片")
|
||||
return cls.decode_base64(encoded)
|
||||
|
||||
@classmethod
|
||||
def decode_base64(cls, encoded: str) -> bytes:
|
||||
"""解码标准或 URL-safe Base64,并确认结果是图片。"""
|
||||
normalized = re.sub(r"\s+", "", encoded).replace("-", "+").replace("_", "/")
|
||||
try:
|
||||
normalized.encode("ascii")
|
||||
except UnicodeEncodeError as exc:
|
||||
raise ValueError("Base64 候选内容含非 ASCII 字符,跳过") from exc
|
||||
if len(normalized) > cls.MAX_IMAGE_BYTES * 4 // 3 + 4:
|
||||
raise ValueError("Base64 图片负载超过 100MB")
|
||||
normalized += "=" * (-len(normalized) % 4)
|
||||
decoded = base64.b64decode(normalized, validate=True)
|
||||
cls.validate_image_bytes(decoded)
|
||||
return decoded
|
||||
|
||||
@staticmethod
|
||||
def validate_image_bytes(data: bytes) -> None:
|
||||
"""拒绝 HTML、JSON 等非图片响应。"""
|
||||
if not data:
|
||||
raise ValueError("响应内容为空")
|
||||
probe = data[:32]
|
||||
is_webp = len(data) >= 12 and data[:4] == b"RIFF" and data[8:12] == b"WEBP"
|
||||
is_avif = (
|
||||
len(data) >= 12
|
||||
and data[4:8] == b"ftyp"
|
||||
and data[8:12] in {b"avif", b"avis"}
|
||||
)
|
||||
is_jp2 = len(data) >= 12 and data[4:8] == b"ftyp" and b"jp2" in data[8:12]
|
||||
signatures = (
|
||||
b"\x89PNG\r\n\x1a\n",
|
||||
b"\xff\xd8\xff",
|
||||
b"GIF87a",
|
||||
b"GIF89a",
|
||||
b"BM",
|
||||
b"II*\x00",
|
||||
b"MM\x00*",
|
||||
b"\x00\x00\x00\x0cjP ",
|
||||
)
|
||||
starts_with_sig = any(probe.find(sig) < 4 for sig in signatures)
|
||||
if starts_with_sig or is_webp or is_avif or is_jp2:
|
||||
return
|
||||
head = data[:64].decode("ascii", errors="ignore").lower()
|
||||
if head.startswith(("<!doctype", "<html", "{", "[")):
|
||||
raise ValueError("响应内容不是图片(检测到 HTML/JSON)")
|
||||
|
||||
async def download_public_image(
|
||||
self, url: str, proxy: str | None = None
|
||||
) -> bytes | None:
|
||||
"""从公网 URL 下载已校验且大小受限的图片。"""
|
||||
try:
|
||||
return await asyncio.wait_for(
|
||||
self._download_image_inner(url, proxy),
|
||||
timeout=self.IMAGE_DOWNLOAD_TOTAL_TIMEOUT,
|
||||
)
|
||||
except asyncio.TimeoutError as exc:
|
||||
raise httpx.TimeoutException(
|
||||
f"图片下载超过 {self.IMAGE_DOWNLOAD_TOTAL_TIMEOUT}s 总超时限制: {self.sanitize_url(url)}"
|
||||
) from exc
|
||||
|
||||
async def _download_image_inner(
|
||||
self, url: str, proxy: str | None = None
|
||||
) -> bytes | None:
|
||||
"""实际下载逻辑。"""
|
||||
current_url = url
|
||||
download_timeout = httpx.Timeout(connect=20.0, read=60.0, write=20.0, pool=20.0)
|
||||
request_proxy = (
|
||||
proxy or self.hooks.config_manager.get_drawing_download_proxy() or None
|
||||
)
|
||||
if request_proxy:
|
||||
logger.debug(
|
||||
"[Comic] 图片下载使用代理: %s", self.sanitize_url(request_proxy)
|
||||
)
|
||||
async with httpx.AsyncClient(
|
||||
timeout=download_timeout,
|
||||
follow_redirects=False,
|
||||
proxy=request_proxy,
|
||||
) as client:
|
||||
for redirect_count in range(self.MAX_IMAGE_REDIRECTS + 1):
|
||||
await self._validate_public_image_url(current_url)
|
||||
logger.info(
|
||||
"[Comic] 正在下载图片 URL: %s", self.sanitize_url(current_url)
|
||||
)
|
||||
resp = await client.get(current_url)
|
||||
if resp.status_code in {301, 302, 303, 307, 308}:
|
||||
location = resp.headers.get("Location")
|
||||
if not location:
|
||||
raise httpx.HTTPStatusError(
|
||||
f"图片重定向缺少地址 [HTTP {resp.status_code}]",
|
||||
request=resp.request,
|
||||
response=resp,
|
||||
)
|
||||
if redirect_count >= self.MAX_IMAGE_REDIRECTS:
|
||||
raise ValueError("图片下载重定向次数超过限制")
|
||||
current_url = str(resp.url.join(location))
|
||||
continue
|
||||
if resp.status_code != 200:
|
||||
raise httpx.HTTPStatusError(
|
||||
f"图片下载失败 [HTTP {resp.status_code}]",
|
||||
request=resp.request,
|
||||
response=resp,
|
||||
)
|
||||
image_bytes = resp.content
|
||||
if len(image_bytes) > self.MAX_IMAGE_BYTES:
|
||||
raise ValueError("图片下载内容超过 100MB")
|
||||
self.validate_image_bytes(image_bytes)
|
||||
return image_bytes
|
||||
raise ValueError("图片下载失败")
|
||||
|
||||
async def _validate_public_image_url(self, url: str) -> None:
|
||||
"""校验图片地址协议与基础合法性,兼容本地与私网自建绘图服务。"""
|
||||
parsed = urlsplit(url)
|
||||
if parsed.scheme not in {"http", "https"} or not parsed.hostname:
|
||||
raise ValueError("图片地址必须是有效的 HTTP/HTTPS URL")
|
||||
if parsed.username or parsed.password:
|
||||
raise ValueError("图片地址不允许包含用户凭据")
|
||||
|
||||
@staticmethod
|
||||
def sanitize_url(url: str) -> str:
|
||||
"""移除日志中的查询参数、片段和用户凭据。"""
|
||||
parsed = urlsplit(url)
|
||||
host = parsed.hostname or ""
|
||||
if parsed.port:
|
||||
host = f"{host}:{parsed.port}"
|
||||
return urlunsplit((parsed.scheme, host, parsed.path, "", ""))
|
||||
|
||||
@staticmethod
|
||||
def summarize_response(data: Any) -> str:
|
||||
"""生成不包含响应正文和 Base64 的结构摘要。"""
|
||||
|
||||
def summarize(value: Any, depth: int = 0) -> str:
|
||||
if isinstance(value, str):
|
||||
return f"<str len={len(value)}>"
|
||||
if depth >= 3:
|
||||
return type(value).__name__
|
||||
if isinstance(value, dict):
|
||||
items = list(value.items())[:10]
|
||||
body = ", ".join(
|
||||
f"{str(key)[:64]}: {summarize(item, depth + 1)}"
|
||||
for key, item in items
|
||||
)
|
||||
suffix = ", ..." if len(value) > len(items) else ""
|
||||
return f"{{{body}{suffix}}}"
|
||||
if isinstance(value, list):
|
||||
items = value[:3]
|
||||
body = ", ".join(summarize(item, depth + 1) for item in items)
|
||||
suffix = ", ..." if len(value) > len(items) else ""
|
||||
return f"[{body}{suffix}] (len={len(value)})"
|
||||
return f"<{type(value).__name__}>"
|
||||
|
||||
return summarize(data)
|
||||
|
||||
|
||||
__all__ = ["DrawingImageResponseService", "ImageDownloadFailedError"]
|
||||
@@ -0,0 +1,266 @@
|
||||
"""
|
||||
群分析插件专用内存日志缓冲与标签提取器
|
||||
提供高性能环形队列日志存储、语义化标签提取与多维度筛选能力。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import re
|
||||
import time
|
||||
from collections import deque
|
||||
from dataclasses import asdict, dataclass
|
||||
from typing import Any
|
||||
|
||||
|
||||
@dataclass
|
||||
class PluginLogEntry:
|
||||
id: str
|
||||
timestamp: float
|
||||
time_str: str
|
||||
level: str
|
||||
logger_name: str
|
||||
trace_id: str | None
|
||||
stage: str | None
|
||||
tag: str
|
||||
message: str
|
||||
raw: str
|
||||
|
||||
def to_dict(self) -> dict[str, Any]:
|
||||
return asdict(self)
|
||||
|
||||
|
||||
class PluginLogBuffer(logging.Handler):
|
||||
"""
|
||||
专用插件日志处理器,挂载到 logging 捕获群分析插件全链路日志
|
||||
"""
|
||||
|
||||
TAG_PATTERNS: list[tuple[str, str, re.Pattern[str]]] = [
|
||||
(
|
||||
"LLM",
|
||||
"大模型调用",
|
||||
re.compile(
|
||||
r"(llm|analyzer|prompt|openai|deepseek|claude|qwen|gpt|token)", re.I
|
||||
),
|
||||
),
|
||||
(
|
||||
"Album",
|
||||
"群相册",
|
||||
re.compile(r"(album|相册|qun_album|group_album)", re.I),
|
||||
),
|
||||
(
|
||||
"OneBot",
|
||||
"OneBot协议",
|
||||
re.compile(r"(onebot|napcat|llonebot|aiocqhttp|gocq)", re.I),
|
||||
),
|
||||
(
|
||||
"QQOfficial",
|
||||
"QQ官方机器人",
|
||||
re.compile(r"(qq_official|botpy|c2c|guild)", re.I),
|
||||
),
|
||||
("Telegram", "Telegram平台", re.compile(r"(telegram|telethon)", re.I)),
|
||||
("Discord", "Discord平台", re.compile(r"(discord|discord_bot)", re.I)),
|
||||
(
|
||||
"Scheduler",
|
||||
"定时与调度",
|
||||
re.compile(r"(scheduler|cron|job|timer|incremental)", re.I),
|
||||
),
|
||||
(
|
||||
"Resilience",
|
||||
"容错与重试",
|
||||
re.compile(r"(resilience|retry|limiter|circuit|lock|reaper)", re.I),
|
||||
),
|
||||
(
|
||||
"Comic",
|
||||
"群漫画",
|
||||
re.compile(
|
||||
r"(comic|漫画|分镜|drawing|storyboard|grok2api|big_banana)", re.I
|
||||
),
|
||||
),
|
||||
(
|
||||
"Render",
|
||||
"报告与长图",
|
||||
re.compile(r"(render|template|html|image|report|playwright|canvas)", re.I),
|
||||
),
|
||||
(
|
||||
"WebUI",
|
||||
"控制台交互",
|
||||
re.compile(r"(webui|bridge|api_|dashboard|route)", re.I),
|
||||
),
|
||||
("Trace", "链路追踪", re.compile(r"(trace|span|context_metric)", re.I)),
|
||||
]
|
||||
|
||||
STAGE_NAMES = {
|
||||
"FETCH_MESSAGES": "拉取聊天记录",
|
||||
"CLEAN_MESSAGES": "消息清洗过滤",
|
||||
"STATS_ANALYSIS": "基础统计分析",
|
||||
"LLM_ANALYSIS": "大模型话题与画像分析",
|
||||
"SAVE_SUMMARY": "历史记录持久化",
|
||||
"RENDER_REPORT": "报告图片渲染与发送",
|
||||
"COMIC_STORYBOARD": "漫画分镜提示词提取",
|
||||
"COMIC_DRAWING": "漫画长图生成与投递",
|
||||
"CRASH_RECOVERY": "异常终止恢复",
|
||||
}
|
||||
|
||||
def __init__(self, max_capacity: int = 2000):
|
||||
super().__init__()
|
||||
self.max_capacity = max_capacity
|
||||
self._buffer: deque[PluginLogEntry] = deque(maxlen=max_capacity)
|
||||
self._counter = 0
|
||||
self._listeners: set[Any] = set()
|
||||
|
||||
def register_listener(self, listener: Any) -> None:
|
||||
"""注册日志实时推送监听器"""
|
||||
self._listeners.add(listener)
|
||||
|
||||
def unregister_listener(self, listener: Any) -> None:
|
||||
"""注销日志实时推送监听器"""
|
||||
self._listeners.discard(listener)
|
||||
|
||||
def record_log(
|
||||
self,
|
||||
level: str,
|
||||
msg: str,
|
||||
trace_id: str | None = None,
|
||||
logger_name: str = "plugin",
|
||||
) -> PluginLogEntry:
|
||||
"""主动记录日志条目并实时推送给前端监听器"""
|
||||
self._counter += 1
|
||||
now = time.time()
|
||||
time_str = time.strftime("%Y-%m-%d %H:%M:%S", time.localtime(now))
|
||||
msecs = int((now - int(now)) * 1000)
|
||||
full_time_str = f"{time_str}.{msecs:03d}"
|
||||
|
||||
# 解析 TraceID:优先从参数取,其次从 TraceContext 上下文取,最后从日志文本中提取
|
||||
if not trace_id:
|
||||
try:
|
||||
from ...shared.trace_context import TraceContext
|
||||
|
||||
ctx = TraceContext.current()
|
||||
if ctx and ctx.trace_id:
|
||||
trace_id = ctx.trace_id
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
if not trace_id:
|
||||
trace_match = re.search(
|
||||
r"\[(manual|incr|group|web_manual|report|[a-zA-Z0-9_\-]+)_[a-zA-Z0-9_\-]+\]",
|
||||
msg,
|
||||
)
|
||||
if trace_match:
|
||||
trace_id = trace_match.group(0).strip("[]")
|
||||
|
||||
# 语义化标签分类
|
||||
tag = "General"
|
||||
for tag_key, _, pattern in self.TAG_PATTERNS:
|
||||
if pattern.search(msg) or pattern.search(logger_name):
|
||||
tag = tag_key
|
||||
break
|
||||
|
||||
# 阶段解析
|
||||
stage = None
|
||||
for stage_code, stage_label in self.STAGE_NAMES.items():
|
||||
if stage_code in msg or stage_label in msg:
|
||||
stage = stage_label
|
||||
break
|
||||
|
||||
entry = PluginLogEntry(
|
||||
id=f"log_{self._counter}",
|
||||
timestamp=now,
|
||||
time_str=full_time_str,
|
||||
level=level.upper(),
|
||||
logger_name=logger_name,
|
||||
trace_id=trace_id,
|
||||
stage=stage,
|
||||
tag=tag,
|
||||
message=msg,
|
||||
raw=f"[{full_time_str}] [{level.upper()}] [{logger_name}]: {msg}",
|
||||
)
|
||||
self._buffer.append(entry)
|
||||
|
||||
for listener in list(self._listeners):
|
||||
try:
|
||||
listener(entry)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
return entry
|
||||
|
||||
def emit(self, record: logging.LogRecord) -> None:
|
||||
try:
|
||||
msg = record.getMessage()
|
||||
logger_name = record.name
|
||||
|
||||
# 仅捕获群分析插件内部日志或带标识日志
|
||||
is_plugin_log = (
|
||||
"astrbot_plugin_qq_group_daily_analysis" in logger_name
|
||||
or "daily_analysis" in logger_name
|
||||
or "[群分析插件]" in msg
|
||||
or hasattr(record, "trace_id")
|
||||
)
|
||||
if not is_plugin_log:
|
||||
return
|
||||
|
||||
trace_id = getattr(record, "trace_id", None)
|
||||
self.record_log(
|
||||
level=record.levelname,
|
||||
msg=msg,
|
||||
trace_id=trace_id,
|
||||
logger_name=logger_name.split(".")[-1]
|
||||
if "." in logger_name
|
||||
else logger_name,
|
||||
)
|
||||
except Exception:
|
||||
self.handleError(record)
|
||||
|
||||
def query(
|
||||
self,
|
||||
limit: int = 100,
|
||||
offset: int = 0,
|
||||
level: str | None = None,
|
||||
trace_id: str | None = None,
|
||||
tag: str | None = None,
|
||||
search: str | None = None,
|
||||
) -> tuple[list[dict[str, Any]], int]:
|
||||
"""按条件多维度筛选日志列表(按时间倒序)"""
|
||||
results: list[PluginLogEntry] = []
|
||||
target_level = level.upper().strip() if level else None
|
||||
target_trace = trace_id.strip() if trace_id else None
|
||||
target_tag = tag.strip() if tag else None
|
||||
search_kw = search.strip().lower() if search else None
|
||||
|
||||
for entry in reversed(self._buffer):
|
||||
if target_level and entry.level != target_level:
|
||||
continue
|
||||
if target_trace and entry.trace_id != target_trace:
|
||||
continue
|
||||
if target_tag and entry.tag.lower() != target_tag.lower():
|
||||
continue
|
||||
if search_kw:
|
||||
if (
|
||||
search_kw not in entry.message.lower()
|
||||
and search_kw not in (entry.trace_id or "").lower()
|
||||
and search_kw not in entry.logger_name.lower()
|
||||
):
|
||||
continue
|
||||
results.append(entry)
|
||||
|
||||
total = len(results)
|
||||
paged = results[offset : offset + limit]
|
||||
return [e.to_dict() for e in paged], total
|
||||
|
||||
def get_trace_logs(self, trace_id: str) -> list[dict[str, Any]]:
|
||||
"""获取特定 TraceID 的全部日志记录"""
|
||||
return [
|
||||
e.to_dict()
|
||||
for e in self._buffer
|
||||
if e.trace_id == trace_id or (trace_id and trace_id in e.message)
|
||||
]
|
||||
|
||||
def clear(self) -> None:
|
||||
"""清空内存缓冲"""
|
||||
self._buffer.clear()
|
||||
|
||||
|
||||
# 全局单例日志缓冲器
|
||||
global_log_buffer = PluginLogBuffer()
|
||||
@@ -0,0 +1,79 @@
|
||||
"""
|
||||
消息发送器 - 基础设施层
|
||||
提供高层消息发送接口,支持跨平台智能路由。
|
||||
"""
|
||||
|
||||
from ...utils.logger import logger
|
||||
|
||||
|
||||
class MessageSender:
|
||||
"""
|
||||
消息发送器
|
||||
封装了 PlatformAdapter 的底层调用,提供更高层的发送接口
|
||||
"""
|
||||
|
||||
def __init__(self, bot_manager, config_manager):
|
||||
self.bot_manager = bot_manager
|
||||
self.config_manager = config_manager
|
||||
|
||||
async def send_text(
|
||||
self, group_id: str, text: str, platform_id: str | None = None
|
||||
) -> bool:
|
||||
"""发送文本消息"""
|
||||
adapter = self.bot_manager.get_adapter(platform_id)
|
||||
if not adapter:
|
||||
logger.error(f"[MessageSender] 未找到平台 {platform_id} 的适配器")
|
||||
return False
|
||||
return await adapter.send_text(group_id, text)
|
||||
|
||||
async def send_image_smart(
|
||||
self,
|
||||
group_id: str,
|
||||
image_url: str,
|
||||
caption: str = "",
|
||||
platform_id: str | None = None,
|
||||
) -> bool:
|
||||
"""智能发送图片,支持自动选择适配器"""
|
||||
adapter = self.bot_manager.get_adapter(platform_id)
|
||||
if not adapter:
|
||||
logger.error(f"[MessageSender] 未找到平台 {platform_id} 的适配器")
|
||||
return False
|
||||
return await adapter.send_image(group_id, image_url, caption)
|
||||
|
||||
async def send_file(
|
||||
self,
|
||||
group_id: str,
|
||||
file_path: str,
|
||||
caption: str = "",
|
||||
platform_id: str | None = None,
|
||||
) -> bool:
|
||||
"""发送文件(HTML/PDF/其它文件)。支持可选 caption。"""
|
||||
adapter = self.bot_manager.get_adapter(platform_id)
|
||||
if not adapter:
|
||||
logger.error(f"[MessageSender] 未找到平台 {platform_id} 的适配器")
|
||||
return False
|
||||
|
||||
# 首先发送文件,本方法的返回值只代表文件是否发送成功。
|
||||
file_sent = await adapter.send_file(group_id, file_path)
|
||||
|
||||
if not file_sent:
|
||||
# 适配器返回 False,表示文件未成功发送
|
||||
return False
|
||||
|
||||
# 文件已成功发送,下面的 caption 发送为尽力而为,不影响整体成功与否。
|
||||
if caption:
|
||||
try:
|
||||
caption_sent = await adapter.send_text(group_id, f"{caption}")
|
||||
if not caption_sent:
|
||||
logger.warning(
|
||||
"[MessageSender] 文件已发送,但 caption 发送失败(适配器返回 False)"
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(f"[MessageSender] 文件已发送,但 caption 发送异常: {e}")
|
||||
|
||||
return True
|
||||
|
||||
def _get_available_platforms(self, group_id: str):
|
||||
"""获取可用的平台列表 (Helper for Dispatcher)"""
|
||||
# 简单实现:返回所有已加载的平台
|
||||
return [(pid, None) for pid in self.bot_manager.get_platform_ids()]
|
||||
@@ -0,0 +1,10 @@
|
||||
"""
|
||||
持久化模块 - 数据存储实现
|
||||
|
||||
包含历史记录仓储和增量分析状态仓储。
|
||||
"""
|
||||
|
||||
from .history_repository import HistoryRepository
|
||||
from .incremental_store import IncrementalStore
|
||||
|
||||
__all__ = ["HistoryRepository", "IncrementalStore"]
|
||||
@@ -0,0 +1,124 @@
|
||||
"""
|
||||
阶段 Checkpoint 缓存仓储 - 支持分析子阶段产物缓存与局部断点续跑 (Partial Resume)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import sqlite3
|
||||
import time
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
|
||||
class CheckpointStore:
|
||||
"""阶段 Checkpoint 存储器,用于在子阶段失败时实现秒级局部重试并节省 Token"""
|
||||
|
||||
def __init__(self, db_path: Path):
|
||||
self.db_path = Path(db_path)
|
||||
self.db_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
self._init_db()
|
||||
|
||||
def _get_connection(self) -> sqlite3.Connection:
|
||||
conn = sqlite3.connect(str(self.db_path), timeout=10.0)
|
||||
conn.row_factory = sqlite3.Row
|
||||
conn.execute("PRAGMA journal_mode=WAL;")
|
||||
return conn
|
||||
|
||||
def _init_db(self) -> None:
|
||||
with self._get_connection() as conn:
|
||||
conn.execute(
|
||||
"""
|
||||
CREATE TABLE IF NOT EXISTS stage_checkpoints (
|
||||
checkpoint_id TEXT PRIMARY KEY,
|
||||
group_id TEXT NOT NULL,
|
||||
date_str TEXT NOT NULL,
|
||||
stage_name TEXT NOT NULL,
|
||||
data_json TEXT NOT NULL,
|
||||
created_at REAL NOT NULL,
|
||||
expire_at REAL NOT NULL
|
||||
);
|
||||
"""
|
||||
)
|
||||
conn.execute(
|
||||
"CREATE INDEX IF NOT EXISTS idx_chk_group_date ON stage_checkpoints(group_id, date_str);"
|
||||
)
|
||||
|
||||
def save_checkpoint(
|
||||
self,
|
||||
group_id: str,
|
||||
date_str: str,
|
||||
stage_name: str,
|
||||
data: Any,
|
||||
ttl_seconds: int = 86400 * 30,
|
||||
) -> None:
|
||||
"""保存阶段产物快照(默认与 Trace 保留期对齐,保留 30 天)"""
|
||||
checkpoint_id = f"{group_id}_{date_str}_{stage_name}"
|
||||
now = time.time()
|
||||
expire_at = now + ttl_seconds
|
||||
|
||||
with self._get_connection() as conn:
|
||||
conn.execute(
|
||||
"""
|
||||
INSERT INTO stage_checkpoints (
|
||||
checkpoint_id, group_id, date_str, stage_name, data_json, created_at, expire_at
|
||||
) VALUES (?, ?, ?, ?, ?, ?, ?)
|
||||
ON CONFLICT(checkpoint_id) DO UPDATE SET
|
||||
data_json=excluded.data_json,
|
||||
created_at=excluded.created_at,
|
||||
expire_at=excluded.expire_at;
|
||||
""",
|
||||
(
|
||||
checkpoint_id,
|
||||
str(group_id),
|
||||
str(date_str),
|
||||
stage_name,
|
||||
json.dumps(data, ensure_ascii=False),
|
||||
now,
|
||||
expire_at,
|
||||
),
|
||||
)
|
||||
|
||||
def get_checkpoint(
|
||||
self, group_id: str, date_str: str, stage_name: str
|
||||
) -> Any | None:
|
||||
"""读取有效的阶段产物快照(若已过期则返回 None 并删除)"""
|
||||
checkpoint_id = f"{group_id}_{date_str}_{stage_name}"
|
||||
now = time.time()
|
||||
|
||||
with self._get_connection() as conn:
|
||||
row = conn.execute(
|
||||
"SELECT * FROM stage_checkpoints WHERE checkpoint_id = ?",
|
||||
(checkpoint_id,),
|
||||
).fetchone()
|
||||
if not row:
|
||||
return None
|
||||
|
||||
if row["expire_at"] < now:
|
||||
conn.execute(
|
||||
"DELETE FROM stage_checkpoints WHERE checkpoint_id = ?",
|
||||
(checkpoint_id,),
|
||||
)
|
||||
return None
|
||||
|
||||
try:
|
||||
return json.loads(row["data_json"])
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
def clear_checkpoints(self, group_id: str, date_str: str) -> None:
|
||||
"""任务全部成功后清理该群当天的临时 Checkpoint"""
|
||||
with self._get_connection() as conn:
|
||||
conn.execute(
|
||||
"DELETE FROM stage_checkpoints WHERE group_id = ? AND date_str = ?",
|
||||
(str(group_id), str(date_str)),
|
||||
)
|
||||
|
||||
def cleanup_expired(self) -> int:
|
||||
"""清理所有已过期的 Checkpoint"""
|
||||
now = time.time()
|
||||
with self._get_connection() as conn:
|
||||
cursor = conn.execute(
|
||||
"DELETE FROM stage_checkpoints WHERE expire_at < ?", (now,)
|
||||
)
|
||||
return cursor.rowcount
|
||||
@@ -0,0 +1,103 @@
|
||||
"""
|
||||
历史记录管理器模块 - 基础设施持久化层
|
||||
负责存储和查询群聊分析报告的摘要信息
|
||||
使用 AstrBot 的 put_kv_data/get_kv_data 实现
|
||||
"""
|
||||
|
||||
import datetime
|
||||
from typing import Any
|
||||
|
||||
from ...utils.logger import logger
|
||||
|
||||
|
||||
class HistoryManager:
|
||||
"""
|
||||
核心组件:历史分析存档管理器
|
||||
|
||||
该类负责将每日生成的群消息分析报告摘要持久化存储,并提供查询接口。
|
||||
底层基于 AstrBot 提供的 KV 存储能力(put_kv_data/get_kv_data),
|
||||
确保即使在 Bot 重启后也能回溯历史数据。
|
||||
"""
|
||||
|
||||
def __init__(self, star_instance: Any):
|
||||
"""
|
||||
初始化历史记录管理器。
|
||||
|
||||
Args:
|
||||
star_instance (Any): Star 插件实例,用于访问底层持久化引擎
|
||||
"""
|
||||
self.plugin = star_instance
|
||||
|
||||
async def save_analysis(
|
||||
self,
|
||||
group_id: str,
|
||||
analysis_result: dict[str, Any],
|
||||
date_str: str | None = None,
|
||||
time_str: str | None = None,
|
||||
) -> bool:
|
||||
"""
|
||||
序列化并存储一份分析报告摘要。
|
||||
|
||||
摘要包含:发言总量、人数、提取的主题摘要及生成时间,不包含完整的原始消息流。
|
||||
|
||||
Args:
|
||||
group_id (str): 群组 ID
|
||||
analysis_result (dict[str, Any]): 包含 statistics, topics, user_titles 的完整分析对象
|
||||
date_str (str, optional): 归档日期 (YYYY-MM-DD),缺省为当天
|
||||
time_str (str, optional): 归档时间点 (HH-MM),缺省为当前时刻
|
||||
|
||||
Returns:
|
||||
bool: 存储是否成功
|
||||
"""
|
||||
try:
|
||||
now = datetime.datetime.now()
|
||||
if not date_str:
|
||||
date_str = now.strftime("%Y-%m-%d")
|
||||
if not time_str:
|
||||
time_str = now.strftime("%H-%M")
|
||||
|
||||
# 消解非法字符,确保 Key 兼容性
|
||||
time_str = time_str.replace(":", "-")
|
||||
|
||||
# 从分析结果中剥离非持久化字段,提取核心统计元数据
|
||||
stats = analysis_result.get("statistics")
|
||||
topics = analysis_result.get("topics", [])
|
||||
user_titles = analysis_result.get("user_titles", [])
|
||||
|
||||
summary = {
|
||||
"message_count": getattr(stats, "message_count", 0) if stats else 0,
|
||||
"participant_count": getattr(stats, "participant_count", 0)
|
||||
if stats
|
||||
else 0,
|
||||
"topics": [{"topic": t.topic, "detail": t.detail} for t in topics],
|
||||
"user_titles_count": len(user_titles),
|
||||
"generated_at": now.strftime("%Y-%m-%d %H:%M:%S"),
|
||||
}
|
||||
|
||||
key = f"analysis_{group_id}_{date_str}_{time_str}"
|
||||
await self.plugin.put_kv_data(key, summary)
|
||||
|
||||
logger.info(
|
||||
f"已保存群 {group_id} 在 {date_str} {time_str} 的分析摘要到历史记录 (Key: {key})"
|
||||
)
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.error(f"保存历史分析记录失败: {e}", exc_info=True)
|
||||
return False
|
||||
|
||||
async def get_history(
|
||||
self, group_id: str, date_str: str, time_str: str
|
||||
) -> dict[str, Any] | None:
|
||||
"""
|
||||
根据群组、日期和时间点检索一份历史摘要。
|
||||
"""
|
||||
time_str = time_str.replace(":", "-")
|
||||
key = f"analysis_{group_id}_{date_str}_{time_str}"
|
||||
return await self.plugin.get_kv_data(key, None)
|
||||
|
||||
async def has_history(self, group_id: str, date_str: str, time_str: str) -> bool:
|
||||
"""
|
||||
快速判定是否存在指定时间点的历史分析记录。
|
||||
"""
|
||||
history = await self.get_history(group_id, date_str, time_str)
|
||||
return history is not None
|
||||
@@ -0,0 +1,212 @@
|
||||
"""
|
||||
历史仓库 - 存储分析历史的实现
|
||||
|
||||
该模块提供分析结果和历史记录的持久化存储。
|
||||
它封装了现有的 history_manager 功能。
|
||||
"""
|
||||
|
||||
import json
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from ...utils.logger import logger
|
||||
|
||||
|
||||
class HistoryRepository:
|
||||
"""
|
||||
基础设施:历史仓库
|
||||
|
||||
负责群聊分析历史记录的持久化存储与检索。目前使用本地 JSON 文件实现,
|
||||
保持了与旧版 `history_manager` 的数据格式兼容性。
|
||||
|
||||
Attributes:
|
||||
data_dir (Path): 插件数据存储的总根目录
|
||||
history_dir (Path): 专门存放历史记录的子目录
|
||||
"""
|
||||
|
||||
def __init__(self, data_dir: str):
|
||||
"""
|
||||
初始化历史仓库。
|
||||
|
||||
Args:
|
||||
data_dir (str): 存储历史数据的基础目录路径
|
||||
"""
|
||||
self.data_dir = Path(data_dir)
|
||||
self.history_dir = self.data_dir / "history"
|
||||
self._ensure_directories()
|
||||
|
||||
def _ensure_directories(self) -> None:
|
||||
"""内部方法:确保所需的目录结构已创建。"""
|
||||
self.history_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
def _get_group_history_path(self, group_id: str) -> Path:
|
||||
"""内部方法:获取特定群组的历史 JSON 文件路径。"""
|
||||
return self.history_dir / f"group_{group_id}.json"
|
||||
|
||||
def save_analysis_result(
|
||||
self,
|
||||
group_id: str,
|
||||
result: dict[str, Any],
|
||||
date_str: str | None = None,
|
||||
) -> bool:
|
||||
"""
|
||||
将分析结果保存到持久化存储。
|
||||
|
||||
Args:
|
||||
group_id (str): 群组标识符
|
||||
result (dict[str, Any]): 包含统计、金句等信息的分析结果字典
|
||||
date_str (str, optional): 关联日期 (YYYY-MM-DD),默认为执行日
|
||||
|
||||
Returns:
|
||||
bool: 保存成功返回 True,发生异常返回 False
|
||||
"""
|
||||
try:
|
||||
date_str = date_str or datetime.now().strftime("%Y-%m-%d")
|
||||
history = self.load_group_history(group_id)
|
||||
|
||||
# 注入执行时间戳
|
||||
if "timestamp" not in result:
|
||||
result["timestamp"] = datetime.now().isoformat()
|
||||
|
||||
# 结构化存储:二级映射 {date -> result}
|
||||
if "daily" not in history:
|
||||
history["daily"] = {}
|
||||
|
||||
history["daily"][date_str] = result
|
||||
history["last_updated"] = datetime.now().isoformat()
|
||||
|
||||
# 原子写入(覆盖)
|
||||
history_path = self._get_group_history_path(group_id)
|
||||
with open(history_path, "w", encoding="utf-8") as f:
|
||||
json.dump(history, f, ensure_ascii=False, indent=2)
|
||||
|
||||
logger.debug(f"已保存群 {group_id} 在 {date_str} 的历史分析记录")
|
||||
return True
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"保存群 {group_id} 的历史记录失败: {e}")
|
||||
return False
|
||||
|
||||
def load_group_history(self, group_id: str) -> dict[str, Any]:
|
||||
"""
|
||||
加载特定群组的完整历史记录字典。
|
||||
|
||||
Args:
|
||||
group_id (str): 群组标识符
|
||||
|
||||
Returns:
|
||||
dict[str, Any]: 历史数据字典,若文件不存在则返回包含空 daily 结构的初始字典
|
||||
"""
|
||||
try:
|
||||
history_path = self._get_group_history_path(group_id)
|
||||
if history_path.exists():
|
||||
with open(history_path, encoding="utf-8") as f:
|
||||
return json.load(f)
|
||||
return {"daily": {}, "group_id": group_id}
|
||||
except Exception as e:
|
||||
logger.error(f"加载群 {group_id} 的历史记录失败: {e}")
|
||||
return {"daily": {}, "group_id": group_id}
|
||||
|
||||
def get_analysis_result(
|
||||
self, group_id: str, date_str: str
|
||||
) -> dict[str, Any] | None:
|
||||
"""
|
||||
获取指定日期已存档的分析结果。
|
||||
|
||||
Args:
|
||||
group_id (str): 群组 ID
|
||||
date_str (str): 目标日期 (YYYY-MM-DD)
|
||||
|
||||
Returns:
|
||||
Optional[dict[str, Any]]: 分析结果字典,未找到则返回 None
|
||||
"""
|
||||
history = self.load_group_history(group_id)
|
||||
return history.get("daily", {}).get(date_str)
|
||||
|
||||
def get_recent_results(self, group_id: str, limit: int = 7) -> list[dict[str, Any]]:
|
||||
"""
|
||||
获取指定群组最近 N 次的分析结果列表。
|
||||
|
||||
Args:
|
||||
group_id (str): 群组 ID
|
||||
limit (int): 最大返回条数
|
||||
|
||||
Returns:
|
||||
list[dict[str, Any]]: 按日期降序排列的结果列表
|
||||
"""
|
||||
history = self.load_group_history(group_id)
|
||||
daily = history.get("daily", {})
|
||||
|
||||
# 按日期字符串字典序降序排列(YYYY-MM-DD 天然有序)
|
||||
sorted_dates = sorted(daily.keys(), reverse=True)[:limit]
|
||||
return [daily[date] for date in sorted_dates]
|
||||
|
||||
def has_analysis_for_date(self, group_id: str, date_str: str) -> bool:
|
||||
"""
|
||||
检查指定日期是否已经生成过分析。
|
||||
|
||||
Args:
|
||||
group_id (str): 群组 ID
|
||||
date_str (str): 日期字符串
|
||||
|
||||
Returns:
|
||||
bool: 存在记录则返回 True
|
||||
"""
|
||||
return self.get_analysis_result(group_id, date_str) is not None
|
||||
|
||||
def delete_old_history(self, group_id: str, keep_days: int = 30) -> int:
|
||||
"""
|
||||
自动清理超过天数限制的陈旧历史记录。
|
||||
|
||||
Args:
|
||||
group_id (str): 群组 ID
|
||||
keep_days (int): 保留的天数上限
|
||||
|
||||
Returns:
|
||||
int: 实际删除的记录条数
|
||||
"""
|
||||
try:
|
||||
history = self.load_group_history(group_id)
|
||||
daily = history.get("daily", {})
|
||||
|
||||
# 计算截止日期边界
|
||||
from datetime import timedelta
|
||||
|
||||
cutoff = (datetime.now() - timedelta(days=keep_days)).strftime("%Y-%m-%d")
|
||||
|
||||
# 筛选已过期的日期
|
||||
dates_to_delete = [date for date in daily.keys() if date < cutoff]
|
||||
|
||||
for date in dates_to_delete:
|
||||
del daily[date]
|
||||
|
||||
if dates_to_delete:
|
||||
history["daily"] = daily
|
||||
history_path = self._get_group_history_path(group_id)
|
||||
with open(history_path, "w", encoding="utf-8") as f:
|
||||
json.dump(history, f, ensure_ascii=False, indent=2)
|
||||
|
||||
return len(dates_to_delete)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"清理群 {group_id} 的陈旧历史记录失败: {e}")
|
||||
return 0
|
||||
|
||||
def list_groups_with_history(self) -> list[str]:
|
||||
"""
|
||||
扫描文件系统,列出当前所有具有存档记录的群组 ID。
|
||||
|
||||
Returns:
|
||||
list[str]: 群组 ID 字符串列表
|
||||
"""
|
||||
try:
|
||||
groups = []
|
||||
for file_path in self.history_dir.glob("group_*.json"):
|
||||
# 从文件名反推群组 ID (group_123.json -> 123)
|
||||
group_id = file_path.stem.replace("group_", "")
|
||||
groups.append(group_id)
|
||||
return groups
|
||||
except Exception as e:
|
||||
logger.error(f"列出历史记录群组失败: {e}")
|
||||
return []
|
||||
@@ -0,0 +1,374 @@
|
||||
"""
|
||||
增量分析批次持久化存储 — 滑动窗口架构
|
||||
|
||||
基于 AstrBot 的 put_kv_data/get_kv_data 实现按批次独立存储,
|
||||
支持按时间窗口查询批次、批次索引管理和过期批次清理。
|
||||
|
||||
KV 键设计:
|
||||
- 批次索引: incr_batch_index_{group_id}
|
||||
值: [{"batch_id": "xxx", "timestamp": 1234567890.0}, ...]
|
||||
- 批次数据: incr_batch_{group_id}_{batch_id}
|
||||
值: IncrementalBatch.to_dict()
|
||||
- 最后分析消息游标: incr_last_ts_{group_id}
|
||||
值: {"timestamp": 1234567890, "message_ids": ["..."]}
|
||||
"""
|
||||
|
||||
from typing import Any
|
||||
|
||||
from ...domain.entities.incremental_state import IncrementalBatch
|
||||
from ...utils.logger import logger
|
||||
|
||||
|
||||
class IncrementalStore:
|
||||
"""
|
||||
增量分析批次持久化仓储
|
||||
|
||||
核心职责:
|
||||
- save_batch: 保存单个批次数据并更新索引
|
||||
- query_batches: 按时间窗口查询批次列表
|
||||
- get_last_analyzed_cursor / update_last_analyzed_cursor: 跨批次去重
|
||||
- cleanup_old_batches: 清理过期批次
|
||||
- get_batch_count: 获取当前批次总数(状态查询用)
|
||||
"""
|
||||
|
||||
# KV 键前缀
|
||||
INDEX_PREFIX = "incr_batch_index"
|
||||
BATCH_PREFIX = "incr_batch"
|
||||
LAST_TS_PREFIX = "incr_last_ts"
|
||||
|
||||
def __init__(self, star_instance: Any):
|
||||
"""
|
||||
初始化批次持久化仓储。
|
||||
|
||||
Args:
|
||||
star_instance: Star 插件实例,用于访问底层 KV 存储引擎
|
||||
"""
|
||||
self.plugin = star_instance
|
||||
|
||||
# ================================================================
|
||||
# 键构建
|
||||
# ================================================================
|
||||
|
||||
def _index_key(self, group_id: str) -> str:
|
||||
"""构建批次索引键"""
|
||||
return f"{self.INDEX_PREFIX}_{group_id}"
|
||||
|
||||
def _batch_key(self, group_id: str, batch_id: str) -> str:
|
||||
"""构建单个批次数据键"""
|
||||
return f"{self.BATCH_PREFIX}_{group_id}_{batch_id}"
|
||||
|
||||
def _last_ts_key(self, group_id: str) -> str:
|
||||
"""构建最后分析消息时间戳键"""
|
||||
return f"{self.LAST_TS_PREFIX}_{group_id}"
|
||||
|
||||
# ================================================================
|
||||
# 批次索引操作
|
||||
# ================================================================
|
||||
|
||||
async def _get_index(self, group_id: str) -> list[dict]:
|
||||
"""
|
||||
获取指定群的批次索引列表。
|
||||
|
||||
Args:
|
||||
group_id: 群组 ID
|
||||
|
||||
Returns:
|
||||
list[dict]: 索引条目列表,每项包含 batch_id 和 timestamp
|
||||
"""
|
||||
key = self._index_key(group_id)
|
||||
try:
|
||||
data = await self.plugin.get_kv_data(key, None)
|
||||
if data is None:
|
||||
return []
|
||||
if isinstance(data, list):
|
||||
return data
|
||||
logger.warning(f"批次索引数据格式异常 (Key: {key}): {type(data)}")
|
||||
return []
|
||||
except Exception as e:
|
||||
logger.error(f"读取批次索引失败 (Key: {key}): {e}", exc_info=True)
|
||||
return []
|
||||
|
||||
async def _save_index(self, group_id: str, index: list[dict]) -> None:
|
||||
"""
|
||||
保存批次索引列表。
|
||||
|
||||
Args:
|
||||
group_id: 群组 ID
|
||||
index: 索引条目列表
|
||||
"""
|
||||
key = self._index_key(group_id)
|
||||
try:
|
||||
await self.plugin.put_kv_data(key, index)
|
||||
except Exception as e:
|
||||
logger.error(f"保存批次索引失败 (Key: {key}): {e}", exc_info=True)
|
||||
raise
|
||||
|
||||
# ================================================================
|
||||
# 批次数据操作
|
||||
# ================================================================
|
||||
|
||||
async def save_batch(self, batch: IncrementalBatch) -> bool:
|
||||
"""
|
||||
保存单个批次数据并更新索引。
|
||||
|
||||
流程:
|
||||
1. 将批次数据写入独立 KV 键
|
||||
2. 将批次元数据(batch_id + timestamp)追加到索引
|
||||
|
||||
Args:
|
||||
batch: 要保存的增量分析批次
|
||||
|
||||
Returns:
|
||||
bool: 保存是否成功
|
||||
"""
|
||||
group_id = batch.group_id
|
||||
batch_key = self._batch_key(group_id, batch.batch_id)
|
||||
|
||||
try:
|
||||
# 1. 保存批次数据
|
||||
await self.plugin.put_kv_data(batch_key, batch.to_dict())
|
||||
|
||||
# 2. 更新索引
|
||||
index = await self._get_index(group_id)
|
||||
existing_entry = next(
|
||||
(entry for entry in index if entry.get("batch_id") == batch.batch_id),
|
||||
None,
|
||||
)
|
||||
if existing_entry is None:
|
||||
index.append(
|
||||
{
|
||||
"batch_id": batch.batch_id,
|
||||
"timestamp": batch.timestamp,
|
||||
}
|
||||
)
|
||||
else:
|
||||
existing_entry["timestamp"] = batch.timestamp
|
||||
await self._save_index(group_id, index)
|
||||
|
||||
logger.debug(
|
||||
f"已保存批次 {batch.batch_id[:8]}... "
|
||||
f"(群 {group_id}, 消息数={batch.messages_count})"
|
||||
)
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"保存批次失败 (群 {group_id}, 批次 {batch.batch_id[:8]}...): {e}",
|
||||
exc_info=True,
|
||||
)
|
||||
return False
|
||||
|
||||
async def query_batches(
|
||||
self,
|
||||
group_id: str,
|
||||
window_start: float,
|
||||
window_end: float,
|
||||
) -> list[IncrementalBatch]:
|
||||
"""
|
||||
按时间窗口查询批次列表。
|
||||
|
||||
从索引中筛选时间戳落在 [window_start, window_end] 范围内的批次,
|
||||
逐个加载完整批次数据。
|
||||
|
||||
Args:
|
||||
group_id: 群组 ID
|
||||
window_start: 窗口起始时间戳(epoch)
|
||||
window_end: 窗口结束时间戳(epoch)
|
||||
|
||||
Returns:
|
||||
list[IncrementalBatch]: 符合窗口范围的批次列表,按时间戳升序
|
||||
"""
|
||||
index = await self._get_index(group_id)
|
||||
|
||||
# 筛选在窗口范围内的批次
|
||||
matching_entries = [
|
||||
entry
|
||||
for entry in index
|
||||
if window_start <= entry.get("timestamp", 0) <= window_end
|
||||
]
|
||||
|
||||
# 按时间戳升序排列
|
||||
matching_entries.sort(key=lambda x: x.get("timestamp", 0))
|
||||
|
||||
batches: list[IncrementalBatch] = []
|
||||
for entry in matching_entries:
|
||||
batch_id = entry.get("batch_id", "")
|
||||
if not batch_id:
|
||||
continue
|
||||
|
||||
batch_key = self._batch_key(group_id, batch_id)
|
||||
try:
|
||||
data = await self.plugin.get_kv_data(batch_key, None)
|
||||
if data is not None:
|
||||
batch = IncrementalBatch.from_dict(data)
|
||||
batches.append(batch)
|
||||
else:
|
||||
logger.warning(
|
||||
f"批次数据缺失 (群 {group_id}, 批次 {batch_id[:8]}...)"
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"加载批次数据失败 (群 {group_id}, 批次 {batch_id[:8]}...): {e}",
|
||||
exc_info=True,
|
||||
)
|
||||
|
||||
logger.debug(
|
||||
f"窗口查询完成: 群 {group_id}, "
|
||||
f"窗口 [{window_start:.0f}, {window_end:.0f}], "
|
||||
f"匹配 {len(batches)}/{len(index)} 个批次"
|
||||
)
|
||||
|
||||
return batches
|
||||
|
||||
# ================================================================
|
||||
# 最后分析消息游标(跨批次去重用)
|
||||
# ================================================================
|
||||
|
||||
async def get_last_analyzed_cursor(self, group_id: str) -> tuple[int, set[str]]:
|
||||
"""获取指定群的最后分析消息游标。
|
||||
|
||||
旧版本仅保存整数时间戳,此处会将其兼容为不含消息 ID 的游标。
|
||||
|
||||
Args:
|
||||
group_id: 群组 ID。
|
||||
|
||||
Returns:
|
||||
最后分析时间戳,以及该时间戳下已经处理的消息 ID 集合。
|
||||
"""
|
||||
key = self._last_ts_key(group_id)
|
||||
try:
|
||||
data = await self.plugin.get_kv_data(key, 0)
|
||||
if isinstance(data, dict):
|
||||
timestamp = max(0, int(data.get("timestamp", 0)))
|
||||
message_ids = data.get("message_ids", [])
|
||||
if not isinstance(message_ids, list):
|
||||
message_ids = []
|
||||
return timestamp, {str(item) for item in message_ids if str(item)}
|
||||
return (int(data) if data else 0), set()
|
||||
except Exception as e:
|
||||
logger.error(f"读取最后分析游标失败 (Key: {key}): {e}", exc_info=True)
|
||||
return 0, set()
|
||||
|
||||
async def update_last_analyzed_cursor(
|
||||
self,
|
||||
group_id: str,
|
||||
timestamp: int,
|
||||
message_ids: set[str],
|
||||
) -> None:
|
||||
"""更新指定群的最后分析消息游标。
|
||||
|
||||
Args:
|
||||
group_id: 群组 ID。
|
||||
timestamp: 最后分析消息的 epoch 时间戳。
|
||||
message_ids: 该时间戳下已经处理的消息 ID。
|
||||
"""
|
||||
key = self._last_ts_key(group_id)
|
||||
try:
|
||||
await self.plugin.put_kv_data(
|
||||
key,
|
||||
{
|
||||
"timestamp": max(0, int(timestamp)),
|
||||
"message_ids": sorted(
|
||||
str(item) for item in message_ids if str(item)
|
||||
),
|
||||
},
|
||||
)
|
||||
logger.debug(f"更新最后分析游标: 群 {group_id}, ts={timestamp}")
|
||||
except Exception as e:
|
||||
logger.error(f"更新最后分析游标失败 (Key: {key}): {e}", exc_info=True)
|
||||
raise
|
||||
|
||||
# ================================================================
|
||||
# 过期批次清理
|
||||
# ================================================================
|
||||
|
||||
async def cleanup_old_batches(self, group_id: str, before_timestamp: float) -> int:
|
||||
"""
|
||||
清理指定群中早于给定时间戳的所有批次。
|
||||
|
||||
流程:
|
||||
1. 从索引中分离出过期条目和保留条目
|
||||
2. 逐个删除过期批次的 KV 数据
|
||||
3. 用保留条目覆盖索引
|
||||
|
||||
Args:
|
||||
group_id: 群组 ID
|
||||
before_timestamp: 清理此时间戳之前的所有批次
|
||||
|
||||
Returns:
|
||||
int: 已清理的批次数量
|
||||
"""
|
||||
index = await self._get_index(group_id)
|
||||
if not index:
|
||||
return 0
|
||||
|
||||
# 分离过期和保留
|
||||
expired = []
|
||||
retained = []
|
||||
for entry in index:
|
||||
if entry.get("timestamp", 0) < before_timestamp:
|
||||
expired.append(entry)
|
||||
else:
|
||||
retained.append(entry)
|
||||
|
||||
if not expired:
|
||||
return 0
|
||||
|
||||
# 删除过期批次数据
|
||||
deleted_count = 0
|
||||
for entry in expired:
|
||||
batch_id = entry.get("batch_id", "")
|
||||
if not batch_id:
|
||||
continue
|
||||
batch_key = self._batch_key(group_id, batch_id)
|
||||
try:
|
||||
await self.plugin.put_kv_data(batch_key, None)
|
||||
deleted_count += 1
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"删除过期批次失败 (群 {group_id}, 批次 {batch_id[:8]}...): {e}",
|
||||
exc_info=True,
|
||||
)
|
||||
|
||||
# 更新索引(仅保留未过期条目)
|
||||
await self._save_index(group_id, retained)
|
||||
|
||||
logger.debug(
|
||||
f"清理过期批次: 群 {group_id}, "
|
||||
f"删除 {deleted_count} 个, 保留 {len(retained)} 个"
|
||||
)
|
||||
|
||||
return deleted_count
|
||||
|
||||
# ================================================================
|
||||
# 状态查询
|
||||
# ================================================================
|
||||
|
||||
async def get_batch_count(self, group_id: str) -> int:
|
||||
"""
|
||||
获取指定群的当前批次总数。
|
||||
|
||||
Args:
|
||||
group_id: 群组 ID
|
||||
|
||||
Returns:
|
||||
int: 批次总数
|
||||
"""
|
||||
index = await self._get_index(group_id)
|
||||
return len(index)
|
||||
|
||||
async def get_all_batch_summaries(self, group_id: str) -> list[dict]:
|
||||
"""
|
||||
获取指定群所有批次的摘要信息(不加载完整数据)。
|
||||
|
||||
用于状态查询命令展示批次概览。
|
||||
|
||||
Args:
|
||||
group_id: 群组 ID
|
||||
|
||||
Returns:
|
||||
list[dict]: 批次摘要列表,按时间升序
|
||||
"""
|
||||
index = await self._get_index(group_id)
|
||||
# 按时间戳升序排列
|
||||
index.sort(key=lambda x: x.get("timestamp", 0))
|
||||
return index
|
||||
@@ -0,0 +1,110 @@
|
||||
"""Persistent registry of groups observed by event-driven platforms."""
|
||||
|
||||
import asyncio
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
|
||||
|
||||
class PlatformGroupRegistry:
|
||||
"""Keep a small, platform-scoped list of groups seen in incoming events."""
|
||||
|
||||
_KV_KEY = "platform_seen_groups_v1"
|
||||
_LEGACY_TELEGRAM_KEY = "telegram_seen_groups_v1"
|
||||
|
||||
def __init__(self, plugin_instance: Any):
|
||||
self.plugin = plugin_instance
|
||||
self._lock = asyncio.Lock()
|
||||
self._known_groups: set[tuple[str, str]] = set()
|
||||
|
||||
async def upsert(
|
||||
self,
|
||||
platform_id: str,
|
||||
group_id: str,
|
||||
sender_id: str = "",
|
||||
sender_name: str = "",
|
||||
event_message_id: str = "",
|
||||
) -> None:
|
||||
platform_key = str(platform_id or "").strip()
|
||||
group_key = str(group_id or "").strip()
|
||||
if not platform_key or not group_key:
|
||||
return
|
||||
|
||||
async with self._lock:
|
||||
identity = (platform_key, group_key)
|
||||
if identity in self._known_groups:
|
||||
return
|
||||
|
||||
registry = await self.plugin.get_kv_data(self._KV_KEY, {})
|
||||
if not isinstance(registry, dict):
|
||||
registry = {}
|
||||
platforms = registry.setdefault("platforms", {})
|
||||
if not isinstance(platforms, dict):
|
||||
platforms = {}
|
||||
registry["platforms"] = platforms
|
||||
platform_map = platforms.setdefault(platform_key, {})
|
||||
if not isinstance(platform_map, dict):
|
||||
platform_map = {}
|
||||
platforms[platform_key] = platform_map
|
||||
|
||||
# Existing groups only need to be remembered in memory. The
|
||||
# registry is used for group discovery, so rewriting last_seen and
|
||||
# the full KV document for every message creates unnecessary I/O.
|
||||
if group_key in platform_map:
|
||||
self._known_groups.add(identity)
|
||||
return
|
||||
|
||||
now_iso = datetime.now(timezone.utc).isoformat()
|
||||
platform_map[group_key] = {
|
||||
"first_seen": now_iso,
|
||||
"last_seen": now_iso,
|
||||
"last_sender_id": str(sender_id or ""),
|
||||
"last_sender_name": str(sender_name or ""),
|
||||
"last_event_message_id": str(event_message_id or ""),
|
||||
}
|
||||
registry["updated_at"] = now_iso
|
||||
await self.plugin.put_kv_data(self._KV_KEY, registry)
|
||||
self._known_groups.add(identity)
|
||||
|
||||
async def get_all_group_ids(self, platform_id: str | None = None) -> list[str]:
|
||||
async with self._lock:
|
||||
registry = await self.plugin.get_kv_data(self._KV_KEY, {})
|
||||
groups = self._extract_groups(registry, platform_id)
|
||||
platform_key = str(platform_id).strip() if platform_id else None
|
||||
if platform_id:
|
||||
self._known_groups.update(
|
||||
(str(platform_key), group_id) for group_id in groups
|
||||
)
|
||||
|
||||
# Preserve groups recorded by older plugin versions.
|
||||
legacy = await self.plugin.get_kv_data(self._LEGACY_TELEGRAM_KEY, {})
|
||||
legacy_groups = self._extract_groups(legacy, platform_id)
|
||||
groups.update(legacy_groups)
|
||||
if platform_id:
|
||||
self._known_groups.update(
|
||||
(str(platform_key), group_id) for group_id in legacy_groups
|
||||
)
|
||||
return sorted(groups)
|
||||
|
||||
@staticmethod
|
||||
def _extract_groups(registry: object, platform_id: str | None) -> set[str]:
|
||||
if not isinstance(registry, dict):
|
||||
return set()
|
||||
platforms = registry.get("platforms")
|
||||
if not isinstance(platforms, dict):
|
||||
return set()
|
||||
|
||||
maps: list[object]
|
||||
if platform_id:
|
||||
maps = [platforms.get(str(platform_id).strip(), {})]
|
||||
else:
|
||||
maps = list(platforms.values())
|
||||
|
||||
groups: set[str] = set()
|
||||
for platform_map in maps:
|
||||
if isinstance(platform_map, dict):
|
||||
groups.update(
|
||||
str(group_id).strip()
|
||||
for group_id in platform_map
|
||||
if str(group_id).strip()
|
||||
)
|
||||
return groups
|
||||
@@ -0,0 +1,102 @@
|
||||
import asyncio
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from astrbot.api.star import Star
|
||||
|
||||
|
||||
class TelegramGroupRegistry:
|
||||
"""
|
||||
Telegram 群组/话题注册表
|
||||
|
||||
负责管理 Telegram 的已见群组和话题列表,用于在无法通过 API 获取群列表时提供回退支持。
|
||||
数据存储在 AstrBot 的 KV 存储中。
|
||||
"""
|
||||
|
||||
_KV_KEY = "telegram_seen_groups_v1"
|
||||
|
||||
def __init__(self, plugin_instance: Star):
|
||||
self.plugin = plugin_instance
|
||||
self._lock = asyncio.Lock()
|
||||
|
||||
async def upsert(
|
||||
self,
|
||||
platform_id: str,
|
||||
group_id: str,
|
||||
sender_id: str,
|
||||
sender_name: str,
|
||||
event_message_id: str,
|
||||
) -> None:
|
||||
"""更新 Telegram 已见群/话题注册表(KV)。"""
|
||||
async with self._lock:
|
||||
registry = await self.plugin.get_kv_data(self._KV_KEY, {})
|
||||
if not isinstance(registry, dict):
|
||||
registry = {}
|
||||
|
||||
platforms = registry.get("platforms")
|
||||
if not isinstance(platforms, dict):
|
||||
platforms = {}
|
||||
registry["platforms"] = platforms
|
||||
|
||||
platform_key = str(platform_id).strip()
|
||||
group_key = str(group_id).strip()
|
||||
|
||||
platform_map = platforms.get(platform_key)
|
||||
if not isinstance(platform_map, dict):
|
||||
platform_map = {}
|
||||
platforms[platform_key] = platform_map
|
||||
|
||||
now_iso = datetime.now(timezone.utc).isoformat()
|
||||
|
||||
entry = platform_map.get(group_key)
|
||||
if not isinstance(entry, dict):
|
||||
entry = {}
|
||||
|
||||
first_seen = entry.get("first_seen")
|
||||
if not isinstance(first_seen, str) or not first_seen:
|
||||
first_seen = now_iso
|
||||
|
||||
entry.update(
|
||||
{
|
||||
"first_seen": first_seen,
|
||||
"last_seen": now_iso,
|
||||
"last_sender_id": str(sender_id),
|
||||
"last_sender_name": str(sender_name),
|
||||
"last_event_message_id": str(event_message_id),
|
||||
}
|
||||
)
|
||||
platform_map[group_key] = entry
|
||||
|
||||
registry["updated_at"] = now_iso
|
||||
await self.plugin.put_kv_data(self._KV_KEY, registry)
|
||||
|
||||
async def get_all_group_ids(self, platform_id: str | None = None) -> list[str]:
|
||||
"""读取 Telegram 已见群/话题列表。"""
|
||||
async with self._lock:
|
||||
registry = await self.plugin.get_kv_data(self._KV_KEY, {})
|
||||
if not isinstance(registry, dict):
|
||||
return []
|
||||
|
||||
platforms = registry.get("platforms")
|
||||
if not isinstance(platforms, dict):
|
||||
return []
|
||||
|
||||
groups: set[str] = set()
|
||||
if platform_id:
|
||||
platform_map = platforms.get(str(platform_id).strip(), {})
|
||||
if isinstance(platform_map, dict):
|
||||
groups.update(
|
||||
str(gid).strip()
|
||||
for gid in platform_map.keys()
|
||||
if str(gid).strip()
|
||||
)
|
||||
else:
|
||||
for platform_map in platforms.values():
|
||||
if not isinstance(platform_map, dict):
|
||||
continue
|
||||
groups.update(
|
||||
str(gid).strip()
|
||||
for gid in platform_map.keys()
|
||||
if str(gid).strip()
|
||||
)
|
||||
|
||||
return sorted(groups)
|
||||
@@ -0,0 +1,902 @@
|
||||
"""
|
||||
Trace 持久化仓储 - 基于 SQLite 的轻量级嵌入式存储
|
||||
负责存储链路快照、细粒度 Span 耗时、dsh-context 风格的上下文演进指标与 Token 消耗审计。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import sqlite3
|
||||
import time
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
|
||||
class TraceSQLiteStore:
|
||||
"""基于 SQLite 的 Trace 链路持久化仓储"""
|
||||
|
||||
def __init__(self, db_path: Path):
|
||||
self.db_path = Path(db_path)
|
||||
self.db_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
self._init_db()
|
||||
|
||||
def _get_connection(self) -> sqlite3.Connection:
|
||||
"""获取启用了 WAL 模式和外键支持的数据库连接"""
|
||||
conn = sqlite3.connect(str(self.db_path), timeout=10.0)
|
||||
conn.row_factory = sqlite3.Row
|
||||
conn.execute("PRAGMA journal_mode=WAL;")
|
||||
conn.execute("PRAGMA foreign_keys=ON;")
|
||||
return conn
|
||||
|
||||
def _init_db(self) -> None:
|
||||
"""初始化数据库表结构与索引"""
|
||||
with self._get_connection() as conn:
|
||||
conn.executescript(
|
||||
"""
|
||||
CREATE TABLE IF NOT EXISTS analysis_traces (
|
||||
trace_id TEXT PRIMARY KEY,
|
||||
group_id TEXT NOT NULL,
|
||||
group_name TEXT DEFAULT '',
|
||||
platform TEXT DEFAULT '',
|
||||
trigger_type TEXT DEFAULT 'manual',
|
||||
status TEXT NOT NULL,
|
||||
started_at REAL NOT NULL,
|
||||
completed_at REAL,
|
||||
duration_ms REAL,
|
||||
error_stage TEXT,
|
||||
error_message TEXT,
|
||||
stack_trace TEXT,
|
||||
extra_json TEXT DEFAULT '{}'
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS trace_spans (
|
||||
span_id TEXT PRIMARY KEY,
|
||||
trace_id TEXT NOT NULL,
|
||||
stage_name TEXT NOT NULL,
|
||||
status TEXT NOT NULL,
|
||||
started_at REAL NOT NULL,
|
||||
duration_ms REAL,
|
||||
stage_payload_json TEXT DEFAULT '{}',
|
||||
FOREIGN KEY (trace_id) REFERENCES analysis_traces(trace_id) ON DELETE CASCADE
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS context_metrics (
|
||||
trace_id TEXT PRIMARY KEY,
|
||||
raw_message_count INTEGER DEFAULT 0,
|
||||
cleaned_message_count INTEGER DEFAULT 0,
|
||||
compression_ratio REAL DEFAULT 0.0,
|
||||
incremental_batches INTEGER DEFAULT 0,
|
||||
window_size INTEGER DEFAULT 0,
|
||||
FOREIGN KEY (trace_id) REFERENCES analysis_traces(trace_id) ON DELETE CASCADE
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS token_usage (
|
||||
trace_id TEXT PRIMARY KEY,
|
||||
prompt_tokens INTEGER DEFAULT 0,
|
||||
completion_tokens INTEGER DEFAULT 0,
|
||||
total_tokens INTEGER DEFAULT 0,
|
||||
estimated_cost REAL DEFAULT 0.0,
|
||||
per_analyzer_tokens_json TEXT DEFAULT '{}',
|
||||
FOREIGN KEY (trace_id) REFERENCES analysis_traces(trace_id) ON DELETE CASCADE
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_traces_started_at ON analysis_traces(started_at DESC);
|
||||
CREATE INDEX IF NOT EXISTS idx_traces_group_id ON analysis_traces(group_id);
|
||||
CREATE INDEX IF NOT EXISTS idx_traces_status ON analysis_traces(status);
|
||||
CREATE INDEX IF NOT EXISTS idx_spans_trace_id ON trace_spans(trace_id);
|
||||
"""
|
||||
)
|
||||
|
||||
def save_trace(self, trace_dict: dict[str, Any]) -> None:
|
||||
"""保存或全量更新 Trace 链路及其关联的 Spans、ContextMetrics、TokenUsage"""
|
||||
trace_id = trace_dict.get("trace_id", "")
|
||||
if not trace_id:
|
||||
return
|
||||
|
||||
with self._get_connection() as conn:
|
||||
# 1. 写入主表前,合并既有的 extra_json(特别是 report_files 产物列表)
|
||||
extra_payload = dict(trace_dict.get("extra", {}))
|
||||
meta = trace_dict.get("metadata", {})
|
||||
if isinstance(meta, dict):
|
||||
for k, v in meta.items():
|
||||
if k not in extra_payload:
|
||||
extra_payload[k] = v
|
||||
elif isinstance(v, dict) and isinstance(extra_payload.get(k), dict):
|
||||
extra_payload[k].update(v)
|
||||
elif isinstance(v, list) and isinstance(extra_payload.get(k), list):
|
||||
extra_payload[k] = v
|
||||
|
||||
existing_row = conn.execute(
|
||||
"SELECT extra_json FROM analysis_traces WHERE trace_id = ?", (trace_id,)
|
||||
).fetchone()
|
||||
if existing_row and existing_row[0]:
|
||||
try:
|
||||
old_extra = json.loads(existing_row[0])
|
||||
merged_rfiles = list(old_extra.get("report_files", []))
|
||||
seen_filenames = {
|
||||
rf.get("filename")
|
||||
for rf in merged_rfiles
|
||||
if isinstance(rf, dict) and rf.get("filename")
|
||||
}
|
||||
for rf in extra_payload.get("report_files", []):
|
||||
if isinstance(rf, dict):
|
||||
fn = rf.get("filename")
|
||||
if fn and fn not in seen_filenames:
|
||||
seen_filenames.add(fn)
|
||||
merged_rfiles.append(rf)
|
||||
extra_payload["report_files"] = merged_rfiles
|
||||
|
||||
# 保留历史已记录的 prompts 与 attempts
|
||||
if (
|
||||
"llm_prompts" not in extra_payload
|
||||
and "llm_prompts" in old_extra
|
||||
):
|
||||
extra_payload["llm_prompts"] = old_extra["llm_prompts"]
|
||||
if (
|
||||
"llm_attempts" not in extra_payload
|
||||
and "llm_attempts" in old_extra
|
||||
):
|
||||
extra_payload["llm_attempts"] = old_extra["llm_attempts"]
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
conn.execute(
|
||||
"""
|
||||
INSERT INTO analysis_traces (
|
||||
trace_id, group_id, group_name, platform, trigger_type,
|
||||
status, started_at, completed_at, duration_ms,
|
||||
error_stage, error_message, stack_trace, extra_json
|
||||
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
ON CONFLICT(trace_id) DO UPDATE SET
|
||||
group_id=CASE WHEN excluded.group_id != '' THEN excluded.group_id ELSE analysis_traces.group_id END,
|
||||
group_name=CASE WHEN excluded.group_name != '' AND excluded.group_name != '未知群' THEN excluded.group_name ELSE analysis_traces.group_name END,
|
||||
platform=CASE WHEN excluded.platform != '' AND excluded.platform NOT IN ('auto', 'default', 'all') THEN excluded.platform ELSE analysis_traces.platform END,
|
||||
status=excluded.status,
|
||||
completed_at=excluded.completed_at,
|
||||
duration_ms=excluded.duration_ms,
|
||||
error_stage=excluded.error_stage,
|
||||
error_message=excluded.error_message,
|
||||
stack_trace=excluded.stack_trace,
|
||||
extra_json=excluded.extra_json;
|
||||
""",
|
||||
(
|
||||
trace_id,
|
||||
str(trace_dict.get("group_id", "")),
|
||||
str(trace_dict.get("group_name", "")),
|
||||
str(trace_dict.get("platform", "")),
|
||||
str(trace_dict.get("trigger_type", "manual")),
|
||||
str(trace_dict.get("status", "running")),
|
||||
float(trace_dict.get("started_at", time.time())),
|
||||
trace_dict.get("completed_at"),
|
||||
trace_dict.get("duration_ms"),
|
||||
trace_dict.get("error_stage"),
|
||||
trace_dict.get("error_message"),
|
||||
trace_dict.get("stack_trace"),
|
||||
json.dumps(extra_payload, ensure_ascii=False),
|
||||
),
|
||||
)
|
||||
|
||||
# 2. 写入 Spans (增量覆写)
|
||||
spans = trace_dict.get("spans", [])
|
||||
for span in spans:
|
||||
conn.execute(
|
||||
"""
|
||||
INSERT INTO trace_spans (
|
||||
span_id, trace_id, stage_name, status, started_at, duration_ms, stage_payload_json
|
||||
) VALUES (?, ?, ?, ?, ?, ?, ?)
|
||||
ON CONFLICT(span_id) DO UPDATE SET
|
||||
status=excluded.status,
|
||||
duration_ms=excluded.duration_ms,
|
||||
stage_payload_json=excluded.stage_payload_json;
|
||||
""",
|
||||
(
|
||||
span.get("span_id", f"{trace_id}_{span.get('stage_name')}"),
|
||||
trace_id,
|
||||
span.get("stage_name", ""),
|
||||
span.get("status", "success"),
|
||||
float(span.get("started_at", time.time())),
|
||||
span.get("duration_ms"),
|
||||
json.dumps(span.get("payload", {}), ensure_ascii=False),
|
||||
),
|
||||
)
|
||||
|
||||
# 3. 写入 Context Metrics
|
||||
context_metrics = trace_dict.get("context_metrics")
|
||||
if context_metrics:
|
||||
conn.execute(
|
||||
"""
|
||||
INSERT INTO context_metrics (
|
||||
trace_id, raw_message_count, cleaned_message_count,
|
||||
compression_ratio, incremental_batches, window_size
|
||||
) VALUES (?, ?, ?, ?, ?, ?)
|
||||
ON CONFLICT(trace_id) DO UPDATE SET
|
||||
raw_message_count=excluded.raw_message_count,
|
||||
cleaned_message_count=excluded.cleaned_message_count,
|
||||
compression_ratio=excluded.compression_ratio,
|
||||
incremental_batches=excluded.incremental_batches,
|
||||
window_size=excluded.window_size;
|
||||
""",
|
||||
(
|
||||
trace_id,
|
||||
int(context_metrics.get("raw_message_count", 0)),
|
||||
int(context_metrics.get("cleaned_message_count", 0)),
|
||||
float(context_metrics.get("compression_ratio", 0.0)),
|
||||
int(context_metrics.get("incremental_batches", 0)),
|
||||
int(context_metrics.get("window_size", 0)),
|
||||
),
|
||||
)
|
||||
|
||||
# 4. 写入 Token Usage
|
||||
token_usage = trace_dict.get("token_usage")
|
||||
if token_usage:
|
||||
conn.execute(
|
||||
"""
|
||||
INSERT INTO token_usage (
|
||||
trace_id, prompt_tokens, completion_tokens, total_tokens,
|
||||
estimated_cost, per_analyzer_tokens_json
|
||||
) VALUES (?, ?, ?, ?, ?, ?)
|
||||
ON CONFLICT(trace_id) DO UPDATE SET
|
||||
prompt_tokens=excluded.prompt_tokens,
|
||||
completion_tokens=excluded.completion_tokens,
|
||||
total_tokens=excluded.total_tokens,
|
||||
estimated_cost=excluded.estimated_cost,
|
||||
per_analyzer_tokens_json=excluded.per_analyzer_tokens_json;
|
||||
""",
|
||||
(
|
||||
trace_id,
|
||||
int(token_usage.get("prompt_tokens", 0)),
|
||||
int(token_usage.get("completion_tokens", 0)),
|
||||
int(token_usage.get("total_tokens", 0)),
|
||||
float(token_usage.get("estimated_cost", 0.0)),
|
||||
json.dumps(
|
||||
token_usage.get("per_analyzer", {}), ensure_ascii=False
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
def get_trace(self, trace_id: str) -> dict[str, Any] | None:
|
||||
"""获取单个 Trace 的完整树状结构(包含 Spans、ContextMetrics、TokenUsage)"""
|
||||
with self._get_connection() as conn:
|
||||
trace_row = conn.execute(
|
||||
"SELECT * FROM analysis_traces WHERE trace_id = ?", (trace_id,)
|
||||
).fetchone()
|
||||
if not trace_row:
|
||||
return None
|
||||
|
||||
trace_data = dict(trace_row)
|
||||
try:
|
||||
trace_data["extra"] = json.loads(trace_data.pop("extra_json") or "{}")
|
||||
except Exception:
|
||||
trace_data["extra"] = {}
|
||||
|
||||
# 查询 Spans
|
||||
spans_rows = conn.execute(
|
||||
"SELECT * FROM trace_spans WHERE trace_id = ? ORDER BY started_at ASC",
|
||||
(trace_id,),
|
||||
).fetchall()
|
||||
spans = []
|
||||
extra_prompts = trace_data.get("extra", {}).get("llm_prompts", {})
|
||||
extra_attempts = trace_data.get("extra", {}).get("llm_attempts", [])
|
||||
for r in spans_rows:
|
||||
s = dict(r)
|
||||
try:
|
||||
s["payload"] = json.loads(s.pop("stage_payload_json") or "{}")
|
||||
except Exception:
|
||||
s["payload"] = {}
|
||||
if s.get("stage_name") == "LLM_ANALYSIS":
|
||||
if extra_prompts and not s["payload"].get("prompts"):
|
||||
s["payload"]["prompts"] = extra_prompts
|
||||
if extra_attempts and not s["payload"].get("llm_attempts"):
|
||||
s["payload"]["llm_attempts"] = extra_attempts
|
||||
spans.append(s)
|
||||
trace_data["spans"] = spans
|
||||
|
||||
# 查询 Context Metrics
|
||||
cm_row = conn.execute(
|
||||
"SELECT * FROM context_metrics WHERE trace_id = ?", (trace_id,)
|
||||
).fetchone()
|
||||
trace_data["context_metrics"] = dict(cm_row) if cm_row else None
|
||||
|
||||
# 查询 Token Usage
|
||||
token_row = conn.execute(
|
||||
"SELECT * FROM token_usage WHERE trace_id = ?", (trace_id,)
|
||||
).fetchone()
|
||||
if token_row:
|
||||
t_data = dict(token_row)
|
||||
try:
|
||||
t_data["per_analyzer"] = json.loads(
|
||||
t_data.pop("per_analyzer_tokens_json") or "{}"
|
||||
)
|
||||
except Exception:
|
||||
t_data["per_analyzer"] = {}
|
||||
trace_data["token_usage"] = t_data
|
||||
else:
|
||||
trace_data["token_usage"] = None
|
||||
|
||||
raw_rfiles = trace_data.get("extra", {}).get("report_files", [])
|
||||
seen_rfiles = set()
|
||||
deduped_rfiles = []
|
||||
for rf in raw_rfiles:
|
||||
fn = rf.get("filename") if isinstance(rf, dict) else None
|
||||
if fn and fn not in seen_rfiles:
|
||||
seen_rfiles.add(fn)
|
||||
deduped_rfiles.append(rf)
|
||||
|
||||
# 额外扫描磁盘 reports 目录,聚合文件名中明确带有 trace_id 的产物文件(如换模板重绘/续跑生成的文件)
|
||||
try:
|
||||
reports_dir = self.db_path.parent / "reports"
|
||||
if reports_dir.is_dir():
|
||||
for p in reports_dir.iterdir():
|
||||
if (
|
||||
p.is_file()
|
||||
and trace_id in p.name
|
||||
and p.name not in seen_rfiles
|
||||
):
|
||||
seen_rfiles.add(p.name)
|
||||
is_html = p.suffix.lower() in (".html", ".htm")
|
||||
is_comic = p.name.lower().startswith(
|
||||
"comic_"
|
||||
) or p.name.startswith("漫画_")
|
||||
stat = p.stat()
|
||||
deduped_rfiles.append(
|
||||
{
|
||||
"filename": p.name,
|
||||
"path": str(p.resolve()),
|
||||
"format": "html" if is_html else "image",
|
||||
"report_type": "comic" if is_comic else "daily",
|
||||
"size_bytes": stat.st_size,
|
||||
"created_at": stat.st_mtime,
|
||||
}
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# 按生成时间降序排序,最新生成的报告排在前面
|
||||
deduped_rfiles.sort(
|
||||
key=lambda x: x.get("created_at", 0) if isinstance(x, dict) else 0,
|
||||
reverse=True,
|
||||
)
|
||||
trace_data["report_files"] = deduped_rfiles
|
||||
return trace_data
|
||||
|
||||
def get_report_trace_map(self) -> dict[str, str]:
|
||||
"""获取已生成的报告文件名与 trace_id 的双向映射"""
|
||||
mapping: dict[str, str] = {}
|
||||
with self._get_connection() as conn:
|
||||
rows = conn.execute(
|
||||
"SELECT trace_id, extra_json FROM analysis_traces WHERE extra_json LIKE '%report_files%'"
|
||||
).fetchall()
|
||||
for r in rows:
|
||||
t_id = str(r["trace_id"])
|
||||
try:
|
||||
extra = json.loads(r["extra_json"] or "{}")
|
||||
for rf in extra.get("report_files", []):
|
||||
fn = rf.get("filename")
|
||||
if fn:
|
||||
mapping[fn] = t_id
|
||||
except Exception:
|
||||
pass
|
||||
return mapping
|
||||
|
||||
def reconcile_crashed_traces_on_startup(self) -> int:
|
||||
"""开机对账扫描:将上次因系统异常终止/重启而未正常收尾的 running 任务标记为 aborted。"""
|
||||
with self._get_connection() as conn:
|
||||
cursor = conn.execute(
|
||||
"""
|
||||
UPDATE analysis_traces
|
||||
SET status = 'aborted',
|
||||
error_stage = 'CRASH_RECOVERY',
|
||||
error_message = 'AstrBot/容器在任务执行期间异常终止,开机已自动回收',
|
||||
completed_at = strftime('%s', 'now')
|
||||
WHERE status = 'running'
|
||||
"""
|
||||
)
|
||||
reconciled_count = cursor.rowcount
|
||||
return reconciled_count
|
||||
|
||||
def get_distinct_groups(self) -> list[dict[str, str]]:
|
||||
"""获取所有有历史分析记录的唯一群组列表(按 group_id 分组,取每个群最新一次运行的群名与平台标识)"""
|
||||
with self._get_connection() as conn:
|
||||
rows = conn.execute(
|
||||
"""
|
||||
WITH ranked_traces AS (
|
||||
SELECT group_id,
|
||||
group_name,
|
||||
platform,
|
||||
started_at,
|
||||
ROW_NUMBER() OVER (PARTITION BY group_id ORDER BY started_at DESC) AS rn
|
||||
FROM analysis_traces
|
||||
WHERE group_id != ''
|
||||
)
|
||||
SELECT r.group_id,
|
||||
COALESCE(
|
||||
NULLIF(r.group_name, ''),
|
||||
NULLIF(r.group_name, '未知群'),
|
||||
(SELECT group_name FROM analysis_traces WHERE group_id = r.group_id AND group_name != '' AND group_name != '未知群' ORDER BY started_at DESC LIMIT 1),
|
||||
r.group_id
|
||||
) AS group_name,
|
||||
r.platform,
|
||||
r.started_at AS last_seen
|
||||
FROM ranked_traces r
|
||||
WHERE r.rn = 1
|
||||
ORDER BY r.started_at DESC;
|
||||
"""
|
||||
).fetchall()
|
||||
return [
|
||||
{
|
||||
"group_id": str(r["group_id"]),
|
||||
"group_name": str(r["group_name"]),
|
||||
"platform": str(r["platform"] or ""),
|
||||
}
|
||||
for r in rows
|
||||
]
|
||||
|
||||
def list_traces(
|
||||
self,
|
||||
limit: int = 20,
|
||||
offset: int = 0,
|
||||
group_id: str | None = None,
|
||||
status: str | None = None,
|
||||
search: str | None = None,
|
||||
start_time: float | None = None,
|
||||
end_time: float | None = None,
|
||||
sort_by: str = "started_at",
|
||||
sort_order: str = "desc",
|
||||
) -> tuple[list[dict[str, Any]], int]:
|
||||
"""分页筛选查询 Trace 列表(支持按群组、状态、关键词、时间范围筛选与排序)"""
|
||||
conditions = []
|
||||
params: list[Any] = []
|
||||
|
||||
if group_id:
|
||||
conditions.append("t.group_id = ?")
|
||||
params.append(str(group_id))
|
||||
if status:
|
||||
conditions.append("t.status = ?")
|
||||
params.append(status)
|
||||
if start_time is not None:
|
||||
conditions.append("t.started_at >= ?")
|
||||
params.append(float(start_time))
|
||||
if end_time is not None:
|
||||
conditions.append("t.started_at <= ?")
|
||||
params.append(float(end_time))
|
||||
if search:
|
||||
conditions.append(
|
||||
"(t.trace_id LIKE ? OR t.group_id LIKE ? OR t.group_name LIKE ?)"
|
||||
)
|
||||
like_pattern = f"%{search}%"
|
||||
params.extend([like_pattern, like_pattern, like_pattern])
|
||||
|
||||
where_clause = f"WHERE {' AND '.join(conditions)}" if conditions else ""
|
||||
|
||||
# 排序字段白名单校验
|
||||
allowed_sort_fields = {
|
||||
"started_at": "t.started_at",
|
||||
"duration_ms": "t.duration_ms",
|
||||
"total_tokens": "tu.total_tokens",
|
||||
"compression_ratio": "cm.compression_ratio",
|
||||
}
|
||||
order_field = allowed_sort_fields.get(sort_by, "t.started_at")
|
||||
order_direction = "ASC" if sort_order.lower() == "asc" else "DESC"
|
||||
|
||||
with self._get_connection() as conn:
|
||||
# 查询总数
|
||||
count_sql = f"SELECT COUNT(*) FROM analysis_traces t {where_clause}"
|
||||
total_count = conn.execute(count_sql, params).fetchone()[0]
|
||||
|
||||
# 查询列表与关联 Token 汇总
|
||||
query_sql = f"""
|
||||
SELECT
|
||||
t.*,
|
||||
tu.total_tokens,
|
||||
tu.estimated_cost,
|
||||
cm.raw_message_count,
|
||||
cm.cleaned_message_count,
|
||||
cm.compression_ratio
|
||||
FROM analysis_traces t
|
||||
LEFT JOIN token_usage tu ON t.trace_id = tu.trace_id
|
||||
LEFT JOIN context_metrics cm ON t.trace_id = cm.trace_id
|
||||
{where_clause}
|
||||
ORDER BY {order_field} {order_direction}
|
||||
LIMIT ? OFFSET ?
|
||||
"""
|
||||
rows = conn.execute(query_sql, params + [limit, offset]).fetchall()
|
||||
|
||||
traces = []
|
||||
for r in rows:
|
||||
item = dict(r)
|
||||
try:
|
||||
item["extra"] = json.loads(item.pop("extra_json") or "{}")
|
||||
except Exception:
|
||||
item["extra"] = {}
|
||||
traces.append(item)
|
||||
|
||||
return traces, total_count
|
||||
|
||||
def get_metrics_summary(self) -> dict[str, Any]:
|
||||
"""获取控制台顶部 KPI 指标与聚合数据"""
|
||||
now = time.time()
|
||||
local_tm = time.localtime(now)
|
||||
# 本地时间今日零点时间戳
|
||||
start_of_today = time.mktime(
|
||||
(local_tm.tm_year, local_tm.tm_mon, local_tm.tm_mday, 0, 0, 0, 0, 0, -1)
|
||||
)
|
||||
|
||||
with self._get_connection() as conn:
|
||||
# 1. 总体概况
|
||||
overview = conn.execute(
|
||||
"""
|
||||
SELECT
|
||||
COUNT(*) as total_traces,
|
||||
SUM(CASE WHEN status IN ('succeeded', 'warning') THEN 1 ELSE 0 END) as succeeded_count,
|
||||
SUM(CASE WHEN status = 'warning' THEN 1 ELSE 0 END) as warning_count,
|
||||
SUM(CASE WHEN status = 'failed' THEN 1 ELSE 0 END) as failed_count,
|
||||
AVG(CASE WHEN status IN ('succeeded', 'warning') THEN duration_ms ELSE NULL END) as avg_duration_ms
|
||||
FROM analysis_traces;
|
||||
"""
|
||||
).fetchone()
|
||||
|
||||
# 2. 今日数据
|
||||
today_data = conn.execute(
|
||||
"""
|
||||
SELECT
|
||||
COUNT(*) as today_traces,
|
||||
COUNT(DISTINCT group_id) as today_active_groups
|
||||
FROM analysis_traces
|
||||
WHERE started_at >= ?;
|
||||
""",
|
||||
(start_of_today,),
|
||||
).fetchone()
|
||||
|
||||
# 3. Token 累计总数与成本
|
||||
token_stats = conn.execute(
|
||||
"""
|
||||
SELECT
|
||||
SUM(total_tokens) as total_tokens_spent,
|
||||
SUM(estimated_cost) as total_cost_spent
|
||||
FROM token_usage;
|
||||
"""
|
||||
).fetchone()
|
||||
|
||||
# 4. 今日 Token 消耗
|
||||
today_tokens = conn.execute(
|
||||
"""
|
||||
SELECT
|
||||
SUM(tu.total_tokens) as today_tokens_spent,
|
||||
SUM(tu.estimated_cost) as today_cost_spent
|
||||
FROM token_usage tu
|
||||
JOIN analysis_traces t ON tu.trace_id = t.trace_id
|
||||
WHERE t.started_at >= ?;
|
||||
""",
|
||||
(start_of_today,),
|
||||
).fetchone()
|
||||
|
||||
return {
|
||||
"total_traces": overview["total_traces"] or 0,
|
||||
"succeeded_count": overview["succeeded_count"] or 0,
|
||||
"failed_count": overview["failed_count"] or 0,
|
||||
"success_rate": (
|
||||
round(
|
||||
(overview["succeeded_count"] or 0)
|
||||
/ max(1, overview["total_traces"] or 1)
|
||||
* 100,
|
||||
1,
|
||||
)
|
||||
),
|
||||
"avg_duration_ms": round(overview["avg_duration_ms"] or 0.0, 1),
|
||||
"today_traces": today_data["today_traces"] or 0,
|
||||
"today_active_groups": today_data["today_active_groups"] or 0,
|
||||
"total_tokens_spent": token_stats["total_tokens_spent"] or 0,
|
||||
"total_cost_spent": round(token_stats["total_cost_spent"] or 0.0, 4),
|
||||
"today_tokens_spent": today_tokens["today_tokens_spent"] or 0,
|
||||
"today_cost_spent": round(today_tokens["today_cost_spent"] or 0.0, 4),
|
||||
"trends": self.get_analytics_trends(granularity="day", range_count=14),
|
||||
}
|
||||
|
||||
def get_analytics_trends(
|
||||
self, granularity: str = "day", range_count: int = 14
|
||||
) -> dict[str, Any]:
|
||||
"""获取时序趋势数据(支持小时 / 天维度切换,并提取服务商与模型消耗细粒度拆分)"""
|
||||
now = time.time()
|
||||
local_tm = time.localtime(now)
|
||||
|
||||
points: list[dict[str, Any]] = []
|
||||
provider_map: dict[str, dict[str, Any]] = {}
|
||||
model_map: dict[str, dict[str, Any]] = {}
|
||||
|
||||
if granularity == "hour":
|
||||
# 按小时划分(默认近 48 小时 / 2 天视野)
|
||||
hours = max(1, min(range_count, 168)) # 上限7天
|
||||
# 当前小时的整点时间戳
|
||||
cur_hour_start = time.mktime(
|
||||
(
|
||||
local_tm.tm_year,
|
||||
local_tm.tm_mon,
|
||||
local_tm.tm_mday,
|
||||
local_tm.tm_hour,
|
||||
0,
|
||||
0,
|
||||
0,
|
||||
0,
|
||||
-1,
|
||||
)
|
||||
)
|
||||
start_timestamp = cur_hour_start - ((hours - 1) * 3600)
|
||||
|
||||
trend_map: dict[str, dict[str, Any]] = {}
|
||||
for i in range(hours):
|
||||
h_ts = start_timestamp + (i * 3600)
|
||||
h_tm = time.localtime(h_ts)
|
||||
h_key = f"{h_tm.tm_year:04d}-{h_tm.tm_mon:02d}-{h_tm.tm_mday:02d} {h_tm.tm_hour:02d}:00"
|
||||
display_date = f"{h_tm.tm_mon}/{h_tm.tm_mday} {h_tm.tm_hour:02d}:00"
|
||||
trend_map[h_key] = {
|
||||
"date": display_date,
|
||||
"date_full": h_key,
|
||||
"timestamp": h_ts,
|
||||
"request_count": 0,
|
||||
"succeeded_count": 0,
|
||||
"failed_count": 0,
|
||||
"prompt_tokens": 0,
|
||||
"completion_tokens": 0,
|
||||
"total_tokens": 0,
|
||||
"estimated_cost": 0.0,
|
||||
}
|
||||
|
||||
with self._get_connection() as conn:
|
||||
rows = conn.execute(
|
||||
"""
|
||||
SELECT
|
||||
strftime('%Y-%m-%d %H:00', datetime(t.started_at, 'unixepoch', 'localtime')) as hour_str,
|
||||
COUNT(t.trace_id) as request_count,
|
||||
SUM(CASE WHEN t.status IN ('succeeded', 'warning') THEN 1 ELSE 0 END) as succeeded_count,
|
||||
SUM(CASE WHEN t.status = 'failed' THEN 1 ELSE 0 END) as failed_count,
|
||||
COALESCE(SUM(tu.prompt_tokens), 0) as prompt_tokens,
|
||||
COALESCE(SUM(tu.completion_tokens), 0) as completion_tokens,
|
||||
COALESCE(SUM(tu.total_tokens), 0) as total_tokens,
|
||||
COALESCE(SUM(tu.estimated_cost), 0.0) as estimated_cost
|
||||
FROM analysis_traces t
|
||||
LEFT JOIN token_usage tu ON t.trace_id = tu.trace_id
|
||||
WHERE t.started_at >= ?
|
||||
GROUP BY hour_str;
|
||||
""",
|
||||
(start_timestamp,),
|
||||
).fetchall()
|
||||
|
||||
for r in rows:
|
||||
h_str = r["hour_str"]
|
||||
if h_str in trend_map:
|
||||
trend_map[h_str]["request_count"] = r["request_count"] or 0
|
||||
trend_map[h_str]["succeeded_count"] = r["succeeded_count"] or 0
|
||||
trend_map[h_str]["failed_count"] = r["failed_count"] or 0
|
||||
trend_map[h_str]["prompt_tokens"] = r["prompt_tokens"] or 0
|
||||
trend_map[h_str]["completion_tokens"] = (
|
||||
r["completion_tokens"] or 0
|
||||
)
|
||||
trend_map[h_str]["total_tokens"] = r["total_tokens"] or 0
|
||||
trend_map[h_str]["estimated_cost"] = round(
|
||||
r["estimated_cost"] or 0.0, 4
|
||||
)
|
||||
|
||||
points = list(trend_map.values())
|
||||
else:
|
||||
# 按天划分(支持 7 天、14 天、30 天等视野)
|
||||
days = max(1, min(range_count, 90))
|
||||
today_start = time.mktime(
|
||||
(
|
||||
local_tm.tm_year,
|
||||
local_tm.tm_mon,
|
||||
local_tm.tm_mday,
|
||||
0,
|
||||
0,
|
||||
0,
|
||||
0,
|
||||
0,
|
||||
-1,
|
||||
)
|
||||
)
|
||||
start_timestamp = today_start - ((days - 1) * 86400)
|
||||
|
||||
trend_map = {}
|
||||
for i in range(days):
|
||||
day_ts = start_timestamp + (i * 86400)
|
||||
d_tm = time.localtime(day_ts)
|
||||
day_key = f"{d_tm.tm_year:04d}-{d_tm.tm_mon:02d}-{d_tm.tm_mday:02d}"
|
||||
display_date = f"{d_tm.tm_mon}/{d_tm.tm_mday}"
|
||||
trend_map[day_key] = {
|
||||
"date": display_date,
|
||||
"date_full": day_key,
|
||||
"timestamp": day_ts,
|
||||
"request_count": 0,
|
||||
"succeeded_count": 0,
|
||||
"failed_count": 0,
|
||||
"prompt_tokens": 0,
|
||||
"completion_tokens": 0,
|
||||
"total_tokens": 0,
|
||||
"estimated_cost": 0.0,
|
||||
}
|
||||
|
||||
with self._get_connection() as conn:
|
||||
rows = conn.execute(
|
||||
"""
|
||||
SELECT
|
||||
date(t.started_at, 'unixepoch', 'localtime') as day_str,
|
||||
COUNT(t.trace_id) as request_count,
|
||||
SUM(CASE WHEN t.status IN ('succeeded', 'warning') THEN 1 ELSE 0 END) as succeeded_count,
|
||||
SUM(CASE WHEN t.status = 'failed' THEN 1 ELSE 0 END) as failed_count,
|
||||
COALESCE(SUM(tu.prompt_tokens), 0) as prompt_tokens,
|
||||
COALESCE(SUM(tu.completion_tokens), 0) as completion_tokens,
|
||||
COALESCE(SUM(tu.total_tokens), 0) as total_tokens,
|
||||
COALESCE(SUM(tu.estimated_cost), 0.0) as estimated_cost
|
||||
FROM analysis_traces t
|
||||
LEFT JOIN token_usage tu ON t.trace_id = tu.trace_id
|
||||
WHERE t.started_at >= ?
|
||||
GROUP BY day_str;
|
||||
""",
|
||||
(start_timestamp,),
|
||||
).fetchall()
|
||||
|
||||
for r in rows:
|
||||
day_str = r["day_str"]
|
||||
if day_str in trend_map:
|
||||
trend_map[day_str]["request_count"] = r["request_count"] or 0
|
||||
trend_map[day_str]["succeeded_count"] = (
|
||||
r["succeeded_count"] or 0
|
||||
)
|
||||
trend_map[day_str]["failed_count"] = r["failed_count"] or 0
|
||||
trend_map[day_str]["prompt_tokens"] = r["prompt_tokens"] or 0
|
||||
trend_map[day_str]["completion_tokens"] = (
|
||||
r["completion_tokens"] or 0
|
||||
)
|
||||
trend_map[day_str]["total_tokens"] = r["total_tokens"] or 0
|
||||
trend_map[day_str]["estimated_cost"] = round(
|
||||
r["estimated_cost"] or 0.0, 4
|
||||
)
|
||||
|
||||
points = list(trend_map.values())
|
||||
|
||||
# 提取该时间范围内的服务商 (Provider) 与模型 (Model) 维度消耗分布
|
||||
with self._get_connection() as conn:
|
||||
traces_with_usage = conn.execute(
|
||||
"""
|
||||
SELECT
|
||||
t.extra_json,
|
||||
COALESCE(tu.total_tokens, 0) as total_tokens,
|
||||
COALESCE(tu.prompt_tokens, 0) as prompt_tokens,
|
||||
COALESCE(tu.completion_tokens, 0) as completion_tokens,
|
||||
COALESCE(tu.per_analyzer_tokens_json, '{}') as per_analyzer
|
||||
FROM analysis_traces t
|
||||
LEFT JOIN token_usage tu ON t.trace_id = tu.trace_id
|
||||
WHERE t.started_at >= ?;
|
||||
""",
|
||||
(start_timestamp,),
|
||||
).fetchall()
|
||||
|
||||
invalid_provider_names = {
|
||||
"default",
|
||||
"default_provider",
|
||||
"quality_provider_id",
|
||||
"none",
|
||||
"",
|
||||
}
|
||||
invalid_model_names = {"default", "default_model", "none", ""}
|
||||
|
||||
for tw in traces_with_usage:
|
||||
try:
|
||||
extra = json.loads(tw["extra_json"] or "{}")
|
||||
except Exception:
|
||||
extra = {}
|
||||
|
||||
llm_prompts = (
|
||||
extra.get("llm_prompts", {}) if isinstance(extra, dict) else {}
|
||||
)
|
||||
|
||||
# 收集有效 provider_id
|
||||
providers_in_trace: set[str] = set()
|
||||
models_in_trace: set[str] = set()
|
||||
|
||||
if isinstance(llm_prompts, dict):
|
||||
for _, p_info in llm_prompts.items():
|
||||
if isinstance(p_info, dict):
|
||||
p_id = str(p_info.get("provider_id", "")).strip()
|
||||
if p_id and p_id.lower() not in invalid_provider_names:
|
||||
providers_in_trace.add(p_id)
|
||||
m_id = str(p_info.get("model", "")).strip()
|
||||
if m_id and m_id.lower() not in invalid_model_names:
|
||||
models_in_trace.add(m_id)
|
||||
|
||||
if not providers_in_trace:
|
||||
fallback_p = str(extra.get("provider_id", "")).strip()
|
||||
if fallback_p and fallback_p.lower() not in invalid_provider_names:
|
||||
providers_in_trace.add(fallback_p)
|
||||
else:
|
||||
providers_in_trace.add("默认会话服务商")
|
||||
|
||||
if not models_in_trace:
|
||||
fallback_m = str(
|
||||
extra.get("model") or extra.get("model_id") or ""
|
||||
).strip()
|
||||
if fallback_m and fallback_m.lower() not in invalid_model_names:
|
||||
models_in_trace.add(fallback_m)
|
||||
else:
|
||||
models_in_trace.add("默认会话模型")
|
||||
|
||||
tot_tok = int(tw["total_tokens"] or 0)
|
||||
tokens_per_p = tot_tok // max(1, len(providers_in_trace))
|
||||
for pid in providers_in_trace:
|
||||
if pid not in provider_map:
|
||||
provider_map[pid] = {
|
||||
"name": pid,
|
||||
"total_tokens": 0,
|
||||
"request_count": 0,
|
||||
}
|
||||
provider_map[pid]["total_tokens"] += tokens_per_p
|
||||
provider_map[pid]["request_count"] += 1
|
||||
|
||||
tokens_per_m = tot_tok // max(1, len(models_in_trace))
|
||||
for mid in models_in_trace:
|
||||
if mid not in model_map:
|
||||
model_map[mid] = {
|
||||
"name": mid,
|
||||
"total_tokens": 0,
|
||||
"request_count": 0,
|
||||
}
|
||||
model_map[mid]["total_tokens"] += tokens_per_m
|
||||
model_map[mid]["request_count"] += 1
|
||||
|
||||
provider_breakdown = [
|
||||
p
|
||||
for p in sorted(
|
||||
provider_map.values(),
|
||||
key=lambda x: x["total_tokens"],
|
||||
reverse=True,
|
||||
)
|
||||
if p["total_tokens"] > 0 or p["request_count"] > 0
|
||||
]
|
||||
model_breakdown = [
|
||||
m
|
||||
for m in sorted(
|
||||
model_map.values(),
|
||||
key=lambda x: x["total_tokens"],
|
||||
reverse=True,
|
||||
)
|
||||
if m["total_tokens"] > 0 or m["request_count"] > 0
|
||||
]
|
||||
|
||||
return {
|
||||
"granularity": granularity,
|
||||
"range_count": range_count,
|
||||
"points": points,
|
||||
"provider_breakdown": provider_breakdown,
|
||||
"model_breakdown": model_breakdown,
|
||||
}
|
||||
|
||||
def cleanup_old_traces(self, days: int = 30, max_count: int = 20000) -> int:
|
||||
"""根据保留天数(默认30天)或最大条数上限清理旧的 Trace 数据,防止 SQLite 无限增长"""
|
||||
cutoff_time = time.time() - (days * 86400)
|
||||
deleted_count = 0
|
||||
|
||||
with self._get_connection() as conn:
|
||||
# 1. 按过期时间清理
|
||||
cursor = conn.execute(
|
||||
"DELETE FROM analysis_traces WHERE started_at < ?", (cutoff_time,)
|
||||
)
|
||||
deleted_count += cursor.rowcount
|
||||
|
||||
# 2. 按最大数量上限清理多余数据
|
||||
total_count = conn.execute(
|
||||
"SELECT COUNT(*) FROM analysis_traces"
|
||||
).fetchone()[0]
|
||||
if total_count > max_count:
|
||||
excess = total_count - max_count
|
||||
conn.execute(
|
||||
"""
|
||||
DELETE FROM analysis_traces
|
||||
WHERE trace_id IN (
|
||||
SELECT trace_id FROM analysis_traces
|
||||
ORDER BY started_at ASC LIMIT ?
|
||||
);
|
||||
""",
|
||||
(excess,),
|
||||
)
|
||||
deleted_count += excess
|
||||
|
||||
return deleted_count
|
||||
@@ -0,0 +1,17 @@
|
||||
from .adapters.discord_adapter import DiscordAdapter
|
||||
from .adapters.lark_adapter import LarkAdapter
|
||||
from .adapters.onebot_adapter import OneBotAdapter
|
||||
from .adapters.qq_official_adapter import QQOfficialAdapter
|
||||
from .adapters.telegram_adapter import TelegramAdapter
|
||||
from .base import PlatformAdapter
|
||||
from .factory import PlatformAdapterFactory
|
||||
|
||||
__all__ = [
|
||||
"PlatformAdapterFactory",
|
||||
"PlatformAdapter",
|
||||
"OneBotAdapter",
|
||||
"LarkAdapter",
|
||||
"QQOfficialAdapter",
|
||||
"TelegramAdapter",
|
||||
"DiscordAdapter",
|
||||
]
|
||||
@@ -0,0 +1,35 @@
|
||||
"""Optional platform adapter exports.
|
||||
|
||||
Each platform is imported independently so an unavailable optional SDK does not
|
||||
prevent the QQ Official adapter from being registered.
|
||||
"""
|
||||
|
||||
__all__: list[str] = []
|
||||
|
||||
try:
|
||||
from .discord_adapter import DiscordAdapter # noqa: F401
|
||||
|
||||
__all__.append("DiscordAdapter")
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
try:
|
||||
from .lark_adapter import LarkAdapter # noqa: F401
|
||||
|
||||
__all__.append("LarkAdapter")
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
try:
|
||||
from .onebot_adapter import OneBotAdapter # noqa: F401
|
||||
|
||||
__all__.append("OneBotAdapter")
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
try:
|
||||
from .qq_official_adapter import QQOfficialAdapter # noqa: F401
|
||||
|
||||
__all__.append("QQOfficialAdapter")
|
||||
except ImportError:
|
||||
pass
|
||||
@@ -0,0 +1,806 @@
|
||||
"""
|
||||
Discord 平台适配器
|
||||
|
||||
为 Discord 平台提供消息获取、发送和群组管理功能。
|
||||
这是一个骨架实现,展示如何为新平台创建适配器。
|
||||
|
||||
注意:Discord 的消息获取需要使用 Discord API,
|
||||
具体实现取决于 AstrBot 的 Discord 集成方式。
|
||||
"""
|
||||
|
||||
from datetime import datetime, timedelta
|
||||
from typing import Any
|
||||
|
||||
from ....utils.logger import logger
|
||||
|
||||
try:
|
||||
import discord
|
||||
except ImportError:
|
||||
discord = None
|
||||
|
||||
from ....domain.value_objects.platform_capabilities import (
|
||||
DISCORD_CAPABILITIES,
|
||||
PlatformCapabilities,
|
||||
)
|
||||
from ....domain.value_objects.unified_group import UnifiedGroup, UnifiedMember
|
||||
from ....domain.value_objects.unified_message import (
|
||||
MessageContent,
|
||||
MessageContentType,
|
||||
UnifiedMessage,
|
||||
)
|
||||
from ..base import PlatformAdapter
|
||||
|
||||
|
||||
class DiscordAdapter(PlatformAdapter):
|
||||
"""
|
||||
具体实现:Discord 平台适配器
|
||||
|
||||
利用 Discord API 为群组(频道)提供消息获取、发送及基础元数据查询功能。
|
||||
由于 Discord 的高度异步特性和复杂的权限模型,该适配器集成了懒加载客户端和多级频道查询机制。
|
||||
|
||||
Attributes:
|
||||
bot_user_id (str): 机器人自身的 Discord 用户 ID
|
||||
"""
|
||||
|
||||
def __init__(self, bot_instance: Any, config: dict | None = None):
|
||||
"""
|
||||
初始化 Discord 适配器。
|
||||
|
||||
Args:
|
||||
bot_instance (Any): 宿主机器人实例
|
||||
config (dict, optional): 配置项,用于提取机器人自身的 Discord ID
|
||||
"""
|
||||
super().__init__(bot_instance, config)
|
||||
# 机器人自己的用户 ID,用于消息过滤(避免分析博取回复)
|
||||
self.bot_user_id = str(config.get("bot_user_id", "")) if config else ""
|
||||
|
||||
# 缓存 Discord 客户端(Lazy Loading)
|
||||
self._cached_client = None
|
||||
|
||||
@property
|
||||
def _discord_client(self) -> Any:
|
||||
"""
|
||||
内部属性:获取实际的 Discord 客户端实例。
|
||||
|
||||
具备懒加载和自动身份嗅探功能。
|
||||
|
||||
Returns:
|
||||
Any: Discord Client 对象
|
||||
"""
|
||||
if self._cached_client:
|
||||
return self._cached_client
|
||||
|
||||
# 执行路径探测逻辑,兼容不同版本的 AstrBot 宿主结构
|
||||
self._cached_client = self._get_discord_client()
|
||||
|
||||
# 兜底:尝试从客户端连接状态中补全机器人 ID
|
||||
if not self.bot_user_id and self._cached_client:
|
||||
if hasattr(self._cached_client, "user") and self._cached_client.user:
|
||||
self.bot_user_id = str(self._cached_client.user.id)
|
||||
|
||||
return self._cached_client
|
||||
|
||||
def _get_discord_client(self) -> Any:
|
||||
"""内部方法:通过多级探测从 bot_instance 中提取 Discord SDK 客户端。"""
|
||||
# 路径 A:bot 本身就是 Client (如小型集成)
|
||||
if hasattr(self.bot, "get_channel"):
|
||||
return self.bot
|
||||
# 路径 B:bot 是包装器,client 在标准成员变量中
|
||||
if hasattr(self.bot, "client"):
|
||||
return self.bot.client
|
||||
# 路径 C:其他常见私有属性名
|
||||
for attr in ("_client", "discord_client", "_discord_client"):
|
||||
if hasattr(self.bot, attr):
|
||||
client = getattr(self.bot, attr)
|
||||
if hasattr(client, "get_channel"):
|
||||
return client
|
||||
logger.warning(f"无法从 {type(self.bot).__name__} 中提取 Discord 客户端实例")
|
||||
return None
|
||||
|
||||
def _init_capabilities(self) -> PlatformCapabilities:
|
||||
"""返回预定义的 Discord 平台能力集。"""
|
||||
return DISCORD_CAPABILITIES
|
||||
|
||||
# ==================== IMessageRepository 实现 ====================
|
||||
|
||||
async def fetch_messages(
|
||||
self,
|
||||
group_id: str,
|
||||
days: int = 1,
|
||||
max_count: int = 100,
|
||||
before_id: str | None = None,
|
||||
since_ts: int | None = None,
|
||||
) -> list[UnifiedMessage]:
|
||||
"""
|
||||
从 Discord 频道异步拉取历史消息记录。
|
||||
|
||||
Args:
|
||||
group_id (str): Discord 频道 (Channel) ID
|
||||
days (int): 查询天数范围
|
||||
max_count (int): 最大拉取消息数量上限
|
||||
before_id (str, optional): 锚点消息 ID,从此之前开始拉取
|
||||
|
||||
Returns:
|
||||
list[UnifiedMessage]: 统一格式的消息对象列表
|
||||
"""
|
||||
if not discord:
|
||||
logger.error("未找到 Discord 模块 (py-cord),无法拉取历史消息。")
|
||||
return []
|
||||
|
||||
try:
|
||||
channel_id = int(group_id)
|
||||
# 先从缓存尝试获取频道
|
||||
channel = self._discord_client.get_channel(channel_id)
|
||||
if not channel:
|
||||
# 缓存未命中则通过网络 fetch
|
||||
try:
|
||||
channel = await self._discord_client.fetch_channel(channel_id)
|
||||
except Exception as e:
|
||||
logger.debug(f"拉取 Discord 频道 {group_id} 失败: {e}")
|
||||
return []
|
||||
|
||||
# 验证权限:确保支持历史消息流
|
||||
if not hasattr(channel, "history"):
|
||||
logger.warning(f"频道 {group_id} 不支持历史消息访问。")
|
||||
return []
|
||||
|
||||
if since_ts and since_ts > 0:
|
||||
start_time = datetime.fromtimestamp(since_ts)
|
||||
else:
|
||||
end_time = datetime.now()
|
||||
start_time = end_time - timedelta(days=days)
|
||||
|
||||
logger.debug(
|
||||
"Discord 消息拉取开始: group=%s, max_count=%s, since_ts=%s, before_id=%s",
|
||||
group_id,
|
||||
max_count,
|
||||
since_ts,
|
||||
before_id,
|
||||
)
|
||||
messages = []
|
||||
|
||||
# 构建 Discord SDK 的 history 查询参数
|
||||
history_kwargs = {"limit": max_count, "after": start_time}
|
||||
if before_id:
|
||||
try:
|
||||
# 使用 Snowflake ID 指向特定消息
|
||||
history_kwargs["before"] = discord.Object(id=int(before_id))
|
||||
except (ValueError, TypeError):
|
||||
pass
|
||||
|
||||
# 消息迭代处理
|
||||
async for msg in channel.history(**history_kwargs):
|
||||
# 排除机器人自身发布的消息
|
||||
if self.bot_user_id and str(msg.author.id) == self.bot_user_id:
|
||||
continue
|
||||
|
||||
unified = self._convert_message(msg, group_id)
|
||||
if unified:
|
||||
messages.append(unified)
|
||||
|
||||
# 排序回升序(SDK 通常返回降序)
|
||||
messages.sort(key=lambda m: m.timestamp)
|
||||
logger.debug(
|
||||
"Discord 消息拉取完成: group=%s, messages=%s",
|
||||
group_id,
|
||||
len(messages),
|
||||
)
|
||||
return messages
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Discord fetch_messages failed: {e}", exc_info=True)
|
||||
return []
|
||||
|
||||
def _convert_message(self, raw_msg: Any, group_id: str) -> UnifiedMessage | None:
|
||||
"""内部方法:将 `discord.Message` 对象转换为统一的 `UnifiedMessage`。"""
|
||||
try:
|
||||
contents = []
|
||||
|
||||
# 1. 基础文本
|
||||
if raw_msg.content:
|
||||
contents.append(
|
||||
MessageContent(type=MessageContentType.TEXT, text=raw_msg.content)
|
||||
)
|
||||
|
||||
# 2. 附件处理 (图片/视频/语音/普通文件)
|
||||
for attachment in raw_msg.attachments:
|
||||
content_type = attachment.content_type or ""
|
||||
if content_type.startswith("image/"):
|
||||
contents.append(
|
||||
MessageContent(
|
||||
type=MessageContentType.IMAGE, url=attachment.url
|
||||
)
|
||||
)
|
||||
elif content_type.startswith("video/"):
|
||||
contents.append(
|
||||
MessageContent(
|
||||
type=MessageContentType.VIDEO, url=attachment.url
|
||||
)
|
||||
)
|
||||
elif content_type.startswith("audio/"):
|
||||
contents.append(
|
||||
MessageContent(
|
||||
type=MessageContentType.VOICE, url=attachment.url
|
||||
)
|
||||
)
|
||||
else:
|
||||
contents.append(
|
||||
MessageContent(
|
||||
type=MessageContentType.FILE,
|
||||
url=attachment.url,
|
||||
raw_data={
|
||||
"filename": attachment.filename,
|
||||
"size": attachment.size,
|
||||
},
|
||||
)
|
||||
)
|
||||
|
||||
# 3. 嵌入内容处理 (部分 Embed 可能包含富文本描述)
|
||||
for embed in raw_msg.embeds:
|
||||
if embed.image:
|
||||
contents.append(
|
||||
MessageContent(
|
||||
type=MessageContentType.IMAGE, url=embed.image.url
|
||||
)
|
||||
)
|
||||
if embed.description:
|
||||
contents.append(
|
||||
MessageContent(
|
||||
type=MessageContentType.TEXT,
|
||||
text=f"\n[Embed] {embed.description}",
|
||||
)
|
||||
)
|
||||
|
||||
# 4. 贴纸处理 (Stickers)
|
||||
if raw_msg.stickers:
|
||||
for sticker in raw_msg.stickers:
|
||||
contents.append(
|
||||
MessageContent(
|
||||
type=MessageContentType.IMAGE, # 贴纸在逻辑上按图片处理
|
||||
url=sticker.url,
|
||||
raw_data={
|
||||
"sticker_id": str(sticker.id),
|
||||
"sticker_name": sticker.name,
|
||||
},
|
||||
)
|
||||
)
|
||||
|
||||
# 确定发送者的显示名称(服务器昵称 > 全局名称 > 用户名)
|
||||
sender_card = None
|
||||
if hasattr(raw_msg.author, "nick") and raw_msg.author.nick:
|
||||
sender_card = raw_msg.author.nick
|
||||
elif hasattr(raw_msg.author, "global_name") and raw_msg.author.global_name:
|
||||
sender_card = raw_msg.author.global_name
|
||||
|
||||
return UnifiedMessage(
|
||||
message_id=str(raw_msg.id),
|
||||
sender_id=str(raw_msg.author.id),
|
||||
sender_name=raw_msg.author.name,
|
||||
sender_card=sender_card,
|
||||
group_id=group_id,
|
||||
text_content=raw_msg.content,
|
||||
contents=tuple(contents),
|
||||
timestamp=int(raw_msg.created_at.timestamp()),
|
||||
platform="discord",
|
||||
reply_to_id=str(raw_msg.reference.message_id)
|
||||
if raw_msg.reference
|
||||
else None,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.debug(f"Discord 消息转换错误: {e}")
|
||||
return None
|
||||
|
||||
def convert_to_raw_format(self, messages: list[UnifiedMessage]) -> list[dict]:
|
||||
"""将统一格式降级转换为 OneBot 风格的字典,以适配下游组件。"""
|
||||
raw_messages = []
|
||||
for msg in messages:
|
||||
raw_msg = {
|
||||
"message_id": msg.message_id,
|
||||
"group_id": msg.group_id,
|
||||
"time": msg.timestamp,
|
||||
"sender": {
|
||||
"user_id": msg.sender_id,
|
||||
"nickname": msg.sender_name,
|
||||
"card": msg.sender_card,
|
||||
},
|
||||
"message": [],
|
||||
"user_id": msg.sender_id, # 后向兼容
|
||||
}
|
||||
|
||||
for content in msg.contents:
|
||||
if content.type == MessageContentType.TEXT:
|
||||
raw_msg["message"].append(
|
||||
{"type": "text", "data": {"text": content.text or ""}}
|
||||
)
|
||||
elif content.type == MessageContentType.IMAGE:
|
||||
raw_msg["message"].append(
|
||||
{
|
||||
"type": "image",
|
||||
"data": {"url": content.url, "file": content.url},
|
||||
}
|
||||
)
|
||||
elif content.type == MessageContentType.AT:
|
||||
raw_msg["message"].append(
|
||||
{"type": "at", "data": {"qq": content.at_user_id}}
|
||||
)
|
||||
elif content.type == MessageContentType.REPLY:
|
||||
if content.raw_data and "reply_id" in content.raw_data:
|
||||
raw_msg["message"].append(
|
||||
{
|
||||
"type": "reply",
|
||||
"data": {"id": content.raw_data["reply_id"]},
|
||||
}
|
||||
)
|
||||
|
||||
raw_messages.append(raw_msg)
|
||||
return raw_messages
|
||||
|
||||
# ==================== IMessageSender 实现 ====================
|
||||
|
||||
async def send_text(
|
||||
self,
|
||||
group_id: str,
|
||||
text: str,
|
||||
reply_to: str | None = None,
|
||||
) -> bool:
|
||||
"""
|
||||
向 Discord 频道发送文本消息。
|
||||
|
||||
Args:
|
||||
group_id (str): 频道 ID
|
||||
text (str): 文本内容
|
||||
reply_to (str, optional): 引用的消息 ID
|
||||
|
||||
Returns:
|
||||
bool: 是否发送成功
|
||||
"""
|
||||
if not discord:
|
||||
return False
|
||||
|
||||
try:
|
||||
channel_id = int(group_id)
|
||||
channel = self.bot.get_channel(channel_id)
|
||||
if not channel:
|
||||
channel = await self.bot.fetch_channel(channel_id)
|
||||
|
||||
if not hasattr(channel, "send"):
|
||||
return False
|
||||
|
||||
reference = None
|
||||
if reply_to:
|
||||
try:
|
||||
reference = discord.MessageReference(
|
||||
message_id=int(reply_to), channel_id=channel_id
|
||||
)
|
||||
except (ValueError, TypeError):
|
||||
pass
|
||||
|
||||
await channel.send(content=text, reference=reference)
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.error(f"Discord 文本发送失败: {e}")
|
||||
return False
|
||||
|
||||
async def send_image(
|
||||
self,
|
||||
group_id: str,
|
||||
image_path: str,
|
||||
caption: str = "",
|
||||
) -> bool:
|
||||
"""
|
||||
向 Discord 频道异步发送图片。
|
||||
|
||||
对于远程 URL,会先下载到内存再通过 Discord API 发送。
|
||||
|
||||
Args:
|
||||
group_id (str): 频道 ID
|
||||
image_path (str): 本地路径或 http URL
|
||||
caption (str): 可选说明文字
|
||||
|
||||
Returns:
|
||||
bool: 是否发送成功
|
||||
"""
|
||||
if not discord:
|
||||
return False
|
||||
|
||||
try:
|
||||
channel_id = int(group_id)
|
||||
channel = self._discord_client.get_channel(channel_id)
|
||||
if not channel:
|
||||
channel = await self._discord_client.fetch_channel(channel_id)
|
||||
|
||||
if not hasattr(channel, "send"):
|
||||
return False
|
||||
|
||||
file_to_send = None
|
||||
if image_path.startswith("base64://"):
|
||||
# Base64 图片:解码 -> 内存 Object -> Discord
|
||||
import base64 # Fix: Ensure base64 is imported
|
||||
from io import BytesIO
|
||||
|
||||
try:
|
||||
base64_data = image_path.split("base64://")[1]
|
||||
image_bytes = base64.b64decode(base64_data)
|
||||
file_to_send = discord.File(
|
||||
BytesIO(image_bytes), filename="daily_report_image.png"
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f"Discord Base64 图片解码失败: {e}")
|
||||
return False
|
||||
|
||||
elif image_path.startswith(("http://", "https://")):
|
||||
# 远程图片:下载 -> 内存 Object -> Discord
|
||||
from io import BytesIO
|
||||
|
||||
import aiohttp
|
||||
|
||||
try:
|
||||
async with aiohttp.ClientSession() as session:
|
||||
async with session.get(
|
||||
image_path, timeout=aiohttp.ClientTimeout(total=30)
|
||||
) as resp:
|
||||
if resp.status == 200:
|
||||
data = await resp.read()
|
||||
# 尽量保留原始后缀
|
||||
filename = image_path.split("/")[-1].split("?")[0]
|
||||
if not filename.lower().endswith(
|
||||
(".png", ".jpg", ".jpeg", ".gif", ".webp")
|
||||
):
|
||||
filename = "daily_report_image.png"
|
||||
|
||||
file_to_send = discord.File(
|
||||
BytesIO(data), filename=filename
|
||||
)
|
||||
else:
|
||||
# 兜底:如果下载失败,直接发 URL 给 Discord 尝试自动解析
|
||||
content = (
|
||||
f"{caption}\n{image_path}"
|
||||
if caption
|
||||
else image_path
|
||||
)
|
||||
await channel.send(content=content)
|
||||
return True
|
||||
except Exception as de:
|
||||
logger.warning(
|
||||
f"Discord 远程图片下载失败: {de},将回退为发送 URL。"
|
||||
)
|
||||
content = f"{caption}\n{image_path}" if caption else image_path
|
||||
await channel.send(content=content)
|
||||
return True
|
||||
else:
|
||||
# 本地图片
|
||||
file_to_send = discord.File(image_path)
|
||||
|
||||
if file_to_send:
|
||||
await channel.send(content=caption or None, file=file_to_send)
|
||||
return True
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Discord 图片发送失败: {e}")
|
||||
return False
|
||||
|
||||
async def send_file(
|
||||
self,
|
||||
group_id: str,
|
||||
file_path: str,
|
||||
filename: str | None = None,
|
||||
) -> bool:
|
||||
"""向 Discord 频道上传任意文件。"""
|
||||
if not discord:
|
||||
return False
|
||||
|
||||
try:
|
||||
channel_id = int(group_id)
|
||||
channel = self.bot.get_channel(channel_id)
|
||||
if not channel:
|
||||
channel = await self.bot.fetch_channel(channel_id)
|
||||
|
||||
if not hasattr(channel, "send"):
|
||||
return False
|
||||
|
||||
file_to_send = discord.File(file_path, filename=filename)
|
||||
await channel.send(file=file_to_send)
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.error(f"Discord 文件发送失败: {e}")
|
||||
return False
|
||||
|
||||
async def send_forward_msg(
|
||||
self,
|
||||
group_id: str,
|
||||
nodes: list[dict],
|
||||
) -> bool:
|
||||
"""
|
||||
在 Discord 模拟合并转发。
|
||||
|
||||
由于 Discord 没有原生节点转发 API,我们将其转换为一组文本消息发送。
|
||||
"""
|
||||
if not discord:
|
||||
return False
|
||||
|
||||
try:
|
||||
channel_id = int(group_id)
|
||||
channel = self._discord_client.get_channel(channel_id)
|
||||
if not channel:
|
||||
channel = await self._discord_client.fetch_channel(channel_id)
|
||||
|
||||
if not hasattr(channel, "send"):
|
||||
return False
|
||||
|
||||
# 将节点汇总为美化的文本块
|
||||
lines = ["📊 **结构化报告摘要 (Structured Report)**\n"]
|
||||
for node in nodes:
|
||||
data = node.get("data", node) # 兼容不同格式
|
||||
name = data.get("name", "AstrBot")
|
||||
content = data.get("content", "")
|
||||
lines.append(f"**[{name}]**:\n{content}\n")
|
||||
|
||||
full_text = "\n".join(lines)
|
||||
|
||||
# 分段处理大消息
|
||||
if len(full_text) > 1900:
|
||||
parts = [
|
||||
full_text[i : i + 1900] for i in range(0, len(full_text), 1900)
|
||||
]
|
||||
for part in parts:
|
||||
await channel.send(content=part)
|
||||
else:
|
||||
await channel.send(content=full_text)
|
||||
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.error(f"Discord 模拟转发失败: {e}")
|
||||
return False
|
||||
|
||||
# ==================== IGroupInfoRepository 实现 ====================
|
||||
|
||||
async def get_group_info(self, group_id: str) -> UnifiedGroup | None:
|
||||
"""解析 Discord 频道及所属服务器的基本信息。"""
|
||||
if not discord:
|
||||
return None
|
||||
|
||||
try:
|
||||
channel_id = int(group_id)
|
||||
channel = self.bot.get_channel(channel_id)
|
||||
if not channel:
|
||||
channel = await self.bot.fetch_channel(channel_id)
|
||||
|
||||
guild = getattr(channel, "guild", None)
|
||||
group_name = getattr(channel, "name", str(channel.id))
|
||||
|
||||
if guild:
|
||||
# 群聊(服务器频道)
|
||||
member_count = guild.member_count
|
||||
owner_id = str(guild.owner_id)
|
||||
else:
|
||||
# 私人对话(DM)
|
||||
member_count = len(getattr(channel, "recipients", [])) + 1
|
||||
owner_id = str(getattr(channel, "owner_id", ""))
|
||||
|
||||
return UnifiedGroup(
|
||||
group_id=str(channel.id),
|
||||
group_name=group_name,
|
||||
member_count=member_count,
|
||||
owner_id=owner_id or None,
|
||||
create_time=int(channel.created_at.timestamp()),
|
||||
platform="discord",
|
||||
)
|
||||
except Exception as e:
|
||||
logger.debug(f"Discord 获取群组信息错误: {e}")
|
||||
return None
|
||||
|
||||
async def get_group_list(self) -> list[str]:
|
||||
"""列出机器人所在服务器中所有可访问的文本频道 ID。"""
|
||||
if not discord:
|
||||
return []
|
||||
|
||||
try:
|
||||
channel_ids = []
|
||||
for guild in self._discord_client.guilds:
|
||||
for channel in guild.text_channels:
|
||||
channel_ids.append(str(channel.id))
|
||||
return channel_ids
|
||||
except Exception:
|
||||
return []
|
||||
|
||||
async def get_member_list(self, group_id: str) -> list[UnifiedMember]:
|
||||
"""
|
||||
获取频道对应的成员列表。
|
||||
|
||||
注意:对于大型服务器,建议启用 GUILD_MEMBERS 意图以保证列表完整性。
|
||||
"""
|
||||
if not discord:
|
||||
return []
|
||||
|
||||
try:
|
||||
channel_id = int(group_id)
|
||||
channel = self.bot.get_channel(channel_id)
|
||||
if not channel:
|
||||
channel = await self.bot.fetch_channel(channel_id)
|
||||
|
||||
guild = getattr(channel, "guild", None)
|
||||
if not guild:
|
||||
# 私聊收件人
|
||||
return [
|
||||
UnifiedMember(
|
||||
user_id=str(u.id),
|
||||
nickname=u.name,
|
||||
card=u.display_name,
|
||||
role="member",
|
||||
)
|
||||
for u in getattr(channel, "recipients", [])
|
||||
]
|
||||
|
||||
members = []
|
||||
for member in guild.members:
|
||||
role = "member"
|
||||
if member.id == guild.owner_id:
|
||||
role = "owner"
|
||||
elif member.guild_permissions.administrator:
|
||||
role = "admin"
|
||||
|
||||
members.append(
|
||||
UnifiedMember(
|
||||
user_id=str(member.id),
|
||||
nickname=member.name,
|
||||
card=member.nick or member.global_name,
|
||||
role=role,
|
||||
join_time=int(member.joined_at.timestamp())
|
||||
if member.joined_at
|
||||
else None,
|
||||
)
|
||||
)
|
||||
return members
|
||||
except Exception:
|
||||
return []
|
||||
|
||||
async def get_member_info(
|
||||
self,
|
||||
group_id: str,
|
||||
user_id: str,
|
||||
) -> UnifiedMember | None:
|
||||
"""获取并解析特定 Discord 用户的身份信息。"""
|
||||
if not discord:
|
||||
return None
|
||||
|
||||
try:
|
||||
uid = int(user_id)
|
||||
channel_id = int(group_id)
|
||||
channel = self.bot.get_channel(channel_id)
|
||||
if not channel:
|
||||
channel = await self.bot.fetch_channel(channel_id)
|
||||
|
||||
guild = getattr(channel, "guild", None)
|
||||
if not guild:
|
||||
# 跨频道/私聊探测
|
||||
user = await self.bot.fetch_user(uid)
|
||||
return UnifiedMember(
|
||||
user_id=str(user.id), nickname=user.name, card=user.display_name
|
||||
)
|
||||
|
||||
member = guild.get_member(uid) or await guild.fetch_member(uid)
|
||||
if not member:
|
||||
return None
|
||||
|
||||
role = (
|
||||
"owner"
|
||||
if member.id == guild.owner_id
|
||||
else ("admin" if member.guild_permissions.administrator else "member")
|
||||
)
|
||||
|
||||
return UnifiedMember(
|
||||
user_id=str(member.id),
|
||||
nickname=member.name,
|
||||
card=member.nick or member.global_name,
|
||||
role=role,
|
||||
join_time=int(member.joined_at.timestamp())
|
||||
if member.joined_at
|
||||
else None,
|
||||
)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
# ==================== IAvatarRepository 实现 ====================
|
||||
|
||||
async def get_user_avatar_url(
|
||||
self,
|
||||
user_id: str,
|
||||
size: int = 100,
|
||||
) -> str | None:
|
||||
"""根据 Discord 用户 ID 动态解析其头像 CDN 地址。"""
|
||||
if not discord or not self._discord_client:
|
||||
return None
|
||||
|
||||
try:
|
||||
uid = int(user_id)
|
||||
user = self._discord_client.get_user(
|
||||
uid
|
||||
) or await self._discord_client.fetch_user(uid)
|
||||
|
||||
if user:
|
||||
# 自动对齐 Discord 支持的尺寸 (2的幂)
|
||||
allowed_sizes = (16, 32, 64, 128, 256, 512, 1024, 2048, 4096)
|
||||
target_size = min(allowed_sizes, key=lambda x: abs(x - size))
|
||||
return user.display_avatar.with_size(target_size).url
|
||||
|
||||
return None
|
||||
except Exception as e:
|
||||
logger.debug(f"Discord 获取用户头像 URL 错误: {e}")
|
||||
return None
|
||||
|
||||
async def get_user_avatar_data(
|
||||
self,
|
||||
user_id: str,
|
||||
size: int = 100,
|
||||
) -> str | None:
|
||||
"""暂不提供 Base64 转换服务,优先使用 CDN 链接。"""
|
||||
return None
|
||||
|
||||
async def get_group_avatar_url(
|
||||
self,
|
||||
group_id: str,
|
||||
size: int = 100,
|
||||
) -> str | None:
|
||||
"""获取 Discord 服务器(Guild)的图标地址。"""
|
||||
if not discord:
|
||||
return None
|
||||
|
||||
try:
|
||||
channel = self.bot.get_channel(
|
||||
int(group_id)
|
||||
) or await self.bot.fetch_channel(int(group_id))
|
||||
guild = getattr(channel, "guild", None)
|
||||
if guild and guild.icon:
|
||||
allowed_sizes = (16, 32, 64, 128, 256, 512, 1024, 2048, 4096)
|
||||
target_size = min(allowed_sizes, key=lambda x: abs(x - size))
|
||||
return guild.icon.with_size(target_size).url
|
||||
return None
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
async def batch_get_avatar_urls(
|
||||
self,
|
||||
user_ids: list[str],
|
||||
size: int = 100,
|
||||
) -> dict[str, str | None]:
|
||||
"""批量获取头像的最佳实践。"""
|
||||
return {uid: await self.get_user_avatar_url(uid, size) for uid in user_ids}
|
||||
|
||||
async def set_reaction(
|
||||
self, group_id: str, message_id: str, emoji: str | int, is_add: bool = True
|
||||
) -> bool:
|
||||
"""
|
||||
Discord 实现消息回应。
|
||||
"""
|
||||
if not discord:
|
||||
return False
|
||||
|
||||
try:
|
||||
reaction_key = str(emoji)
|
||||
emoji_to_use = {
|
||||
"analysis_started": "🔍",
|
||||
"analysis_done": "📊",
|
||||
"289": "🔍",
|
||||
"124": "📊",
|
||||
"424": "📊",
|
||||
}.get(reaction_key, reaction_key)
|
||||
|
||||
channel_id = int(group_id)
|
||||
channel = self._discord_client.get_channel(channel_id)
|
||||
if not channel:
|
||||
channel = await self._discord_client.fetch_channel(channel_id)
|
||||
|
||||
if not hasattr(channel, "get_partial_message"):
|
||||
# 如果较低版本的 SDK 没这个方法,则直接 fetch
|
||||
msg = await channel.fetch_message(int(message_id))
|
||||
else:
|
||||
msg = channel.get_partial_message(int(message_id))
|
||||
|
||||
if is_add:
|
||||
await msg.add_reaction(emoji_to_use)
|
||||
else:
|
||||
await msg.remove_reaction(emoji_to_use, self._discord_client.user)
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.debug(f"Discord set_reaction 失败: {e}")
|
||||
return False
|
||||
@@ -0,0 +1,582 @@
|
||||
"""QQ Official Bot adapter backed by AstrBot's local message history."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import base64
|
||||
import hashlib
|
||||
import os
|
||||
import random
|
||||
import re
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from urllib.parse import quote
|
||||
|
||||
import aiohttp
|
||||
|
||||
from ....domain.value_objects.platform_capabilities import (
|
||||
QQ_OFFICIAL_CAPABILITIES,
|
||||
PlatformCapabilities,
|
||||
)
|
||||
from ....domain.value_objects.unified_group import UnifiedGroup, UnifiedMember
|
||||
from ....domain.value_objects.unified_message import (
|
||||
MessageContent,
|
||||
MessageContentType,
|
||||
UnifiedMessage,
|
||||
)
|
||||
from ....utils.logger import logger
|
||||
from ..base import PlatformAdapter
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from astrbot.api.star import Context
|
||||
|
||||
|
||||
class QQOfficialAdapter(PlatformAdapter):
|
||||
"""Adapter for QQ Official group bots (WebSocket and Webhook variants)."""
|
||||
|
||||
platform_name = "qq_official"
|
||||
AVATAR_TEMPLATE = "https://thirdqq.qlogo.cn/qqapp/{appid}/{member_openid}/640"
|
||||
HISTORY_PAGE_SIZE = 500
|
||||
MARKDOWN_CHUNK_SIZE = 3900
|
||||
|
||||
def __init__(self, bot_instance: Any, config: dict | None = None):
|
||||
super().__init__(bot_instance, config)
|
||||
self._context: Context | None = None
|
||||
self._plugin_instance = config.get("plugin_instance") if config else None
|
||||
self._platform_id = str(config.get("platform_id", "")).strip() if config else ""
|
||||
ids = config.get("bot_self_ids", []) if config else []
|
||||
self.bot_self_ids = [str(item) for item in ids if item]
|
||||
self.appid = self._resolve_appid(config or {})
|
||||
self._markdown_msg_seq = random.randint(1, 10000)
|
||||
self._member_profiles: dict[str, dict[str, str]] = {}
|
||||
|
||||
@property
|
||||
def platform_id(self) -> str:
|
||||
return self._platform_id or "qq_official"
|
||||
|
||||
def _resolve_appid(self, config: dict) -> str:
|
||||
direct = str(config.get("appid", "") or "").strip()
|
||||
if direct:
|
||||
return direct
|
||||
platform = getattr(self.bot, "platform", None)
|
||||
platform_config = getattr(platform, "config", None)
|
||||
if isinstance(platform_config, dict):
|
||||
return str(platform_config.get("appid", "") or "").strip()
|
||||
return ""
|
||||
|
||||
@staticmethod
|
||||
def _is_placeholder_sender_name(name: str | None, sender_id: str) -> bool:
|
||||
normalized = str(name or "").strip()
|
||||
if not normalized:
|
||||
return True
|
||||
if normalized.lower() in {"unknown", "none", "null", "nil", "undefined"}:
|
||||
return True
|
||||
return normalized == str(sender_id).strip()
|
||||
|
||||
@classmethod
|
||||
def _resolve_history_sender_name(
|
||||
cls, sender_name: str | None, sender_id: str, group_id: str
|
||||
) -> str:
|
||||
normalized = str(sender_name or "").strip()
|
||||
if not cls._is_placeholder_sender_name(normalized, sender_id):
|
||||
return normalized
|
||||
digest = hashlib.sha256(f"{group_id}\0{sender_id}".encode()).hexdigest()[:8]
|
||||
return f"群友-{digest.upper()}"
|
||||
|
||||
def set_context(self, context: Context) -> None:
|
||||
self._context = context
|
||||
|
||||
def remember_user_profile(
|
||||
self, user_id: str, nickname: str = "", avatar_url: str = ""
|
||||
) -> None:
|
||||
"""缓存 QQ 官方消息事件提供的用户资料。
|
||||
|
||||
Args:
|
||||
user_id: QQ 成员 OpenID。
|
||||
nickname: 事件提供的显示名称。
|
||||
avatar_url: 事件提供的头像地址。
|
||||
"""
|
||||
normalized_user_id = str(user_id or "").strip()
|
||||
if not normalized_user_id:
|
||||
return
|
||||
|
||||
profile = self._member_profiles.setdefault(normalized_user_id, {})
|
||||
normalized_nickname = str(nickname or "").strip()
|
||||
if not self._is_placeholder_sender_name(
|
||||
normalized_nickname, normalized_user_id
|
||||
):
|
||||
profile["nickname"] = normalized_nickname
|
||||
|
||||
normalized_avatar_url = str(avatar_url or "").strip()
|
||||
if normalized_avatar_url.startswith(("https://", "http://")):
|
||||
profile["avatar_url"] = normalized_avatar_url
|
||||
|
||||
def _init_capabilities(self) -> PlatformCapabilities:
|
||||
return QQ_OFFICIAL_CAPABILITIES
|
||||
|
||||
async def fetch_messages(
|
||||
self,
|
||||
group_id: str,
|
||||
days: int = 1,
|
||||
max_count: int = 1000,
|
||||
before_id: str | None = None,
|
||||
since_ts: int | None = None,
|
||||
) -> list[UnifiedMessage]:
|
||||
if not self._context:
|
||||
logger.warning("[QQOfficial] 未设置 context,无法读取本地消息历史")
|
||||
return []
|
||||
|
||||
history_mgr = self._context.message_history_manager
|
||||
target_count = max(1, int(max_count))
|
||||
cutoff_ts = (
|
||||
int(since_ts)
|
||||
if since_ts and since_ts > 0
|
||||
else int((datetime.now(timezone.utc) - timedelta(days=days)).timestamp())
|
||||
)
|
||||
before_record_id: int | None = None
|
||||
if before_id:
|
||||
try:
|
||||
before_record_id = int(before_id)
|
||||
except (TypeError, ValueError):
|
||||
pass
|
||||
|
||||
messages: list[UnifiedMessage] = []
|
||||
seen_message_ids: set[str] = set()
|
||||
page = 1
|
||||
|
||||
try:
|
||||
while len(messages) < target_count:
|
||||
records = await history_mgr.get(
|
||||
platform_id=self.platform_id,
|
||||
user_id=str(group_id),
|
||||
page=page,
|
||||
page_size=self.HISTORY_PAGE_SIZE,
|
||||
)
|
||||
if not records:
|
||||
break
|
||||
|
||||
reached_cutoff = False
|
||||
for record in records:
|
||||
record_id = getattr(record, "id", None)
|
||||
if (
|
||||
before_record_id is not None
|
||||
and record_id is not None
|
||||
and int(record_id) >= before_record_id
|
||||
):
|
||||
continue
|
||||
|
||||
unified = self._convert_history_record(record, str(group_id))
|
||||
if not unified:
|
||||
continue
|
||||
if unified.timestamp < cutoff_ts:
|
||||
reached_cutoff = True
|
||||
continue
|
||||
if unified.sender_id in self.bot_self_ids:
|
||||
continue
|
||||
if unified.message_id in seen_message_ids:
|
||||
continue
|
||||
|
||||
seen_message_ids.add(unified.message_id)
|
||||
messages.append(unified)
|
||||
|
||||
if len(messages) >= target_count:
|
||||
break
|
||||
if reached_cutoff or len(records) < self.HISTORY_PAGE_SIZE:
|
||||
break
|
||||
page += 1
|
||||
|
||||
messages.sort(key=lambda item: (item.timestamp, item.message_id))
|
||||
if len(messages) > target_count:
|
||||
messages = messages[-target_count:]
|
||||
logger.info(
|
||||
"[QQOfficial] 从本地历史获取群 %s 消息 %s 条",
|
||||
group_id,
|
||||
len(messages),
|
||||
)
|
||||
return messages
|
||||
except Exception as exc:
|
||||
logger.error("[QQOfficial] 读取本地消息历史失败: %s", exc, exc_info=True)
|
||||
return []
|
||||
|
||||
def _convert_history_record(
|
||||
self, record: Any, group_id: str
|
||||
) -> UnifiedMessage | None:
|
||||
try:
|
||||
content = getattr(record, "content", None)
|
||||
if not isinstance(content, dict):
|
||||
return None
|
||||
metadata = content.get("_qq_official")
|
||||
if not isinstance(metadata, dict):
|
||||
return None
|
||||
|
||||
contents: list[MessageContent] = []
|
||||
text_parts: list[str] = []
|
||||
for part in content.get("message", []):
|
||||
if not isinstance(part, dict):
|
||||
continue
|
||||
part_type = str(part.get("type", "")).lower()
|
||||
if part_type in {"plain", "text"}:
|
||||
text = str(part.get("text", "") or "")
|
||||
text_parts.append(text)
|
||||
contents.append(
|
||||
MessageContent(type=MessageContentType.TEXT, text=text)
|
||||
)
|
||||
elif part_type == "image":
|
||||
contents.append(
|
||||
MessageContent(
|
||||
type=MessageContentType.IMAGE,
|
||||
url=str(part.get("url", "") or ""),
|
||||
)
|
||||
)
|
||||
elif part_type == "at":
|
||||
contents.append(
|
||||
MessageContent(
|
||||
type=MessageContentType.AT,
|
||||
at_user_id=str(part.get("target_id", "") or ""),
|
||||
)
|
||||
)
|
||||
elif part_type == "file":
|
||||
contents.append(
|
||||
MessageContent(
|
||||
type=MessageContentType.FILE,
|
||||
url=str(part.get("url", "") or ""),
|
||||
raw_data={"name": part.get("name", "")},
|
||||
)
|
||||
)
|
||||
elif part_type in {"record", "voice"}:
|
||||
contents.append(
|
||||
MessageContent(
|
||||
type=MessageContentType.VOICE,
|
||||
url=str(part.get("url", "") or ""),
|
||||
)
|
||||
)
|
||||
elif part_type == "video":
|
||||
contents.append(
|
||||
MessageContent(
|
||||
type=MessageContentType.VIDEO,
|
||||
url=str(part.get("url", "") or ""),
|
||||
)
|
||||
)
|
||||
|
||||
message_id = str(metadata.get("message_id", "") or "")
|
||||
if not message_id:
|
||||
message_id = f"local:{getattr(record, 'id', '')}"
|
||||
timestamp = int(metadata.get("timestamp", 0) or 0)
|
||||
if timestamp <= 0:
|
||||
created_at = getattr(record, "created_at", None)
|
||||
timestamp = int(created_at.timestamp()) if created_at else 0
|
||||
sender_id = str(getattr(record, "sender_id", "") or "")
|
||||
if not sender_id:
|
||||
return None
|
||||
sender_name = self._resolve_history_sender_name(
|
||||
getattr(record, "sender_name", None), sender_id, group_id
|
||||
)
|
||||
|
||||
return UnifiedMessage(
|
||||
message_id=message_id,
|
||||
sender_id=sender_id,
|
||||
sender_name=sender_name,
|
||||
sender_card=None,
|
||||
group_id=group_id,
|
||||
text_content="".join(text_parts),
|
||||
contents=tuple(contents),
|
||||
timestamp=timestamp,
|
||||
platform=self.platform_name,
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.debug("[QQOfficial] 转换本地历史记录失败: %s", exc)
|
||||
return None
|
||||
|
||||
def convert_to_raw_format(self, messages: list[UnifiedMessage]) -> list[dict]:
|
||||
result: list[dict] = []
|
||||
for message in messages:
|
||||
chain: list[dict] = []
|
||||
for content in message.contents:
|
||||
if content.type == MessageContentType.TEXT:
|
||||
chain.append({"type": "text", "data": {"text": content.text}})
|
||||
elif content.type == MessageContentType.IMAGE:
|
||||
chain.append({"type": "image", "data": {"url": content.url}})
|
||||
elif content.type == MessageContentType.AT:
|
||||
chain.append({"type": "at", "data": {"qq": content.at_user_id}})
|
||||
result.append(
|
||||
{
|
||||
"message_id": message.message_id,
|
||||
"time": message.timestamp,
|
||||
"group_id": message.group_id,
|
||||
"sender": {
|
||||
"user_id": message.sender_id,
|
||||
"nickname": message.sender_name or message.sender_id,
|
||||
"card": "",
|
||||
},
|
||||
"message": chain,
|
||||
"user_id": message.sender_id,
|
||||
}
|
||||
)
|
||||
return result
|
||||
|
||||
async def _send_chain(self, group_id: str, chain: Any) -> bool:
|
||||
if not self._context:
|
||||
logger.error("[QQOfficial] 未设置 context,无法发送消息")
|
||||
return False
|
||||
try:
|
||||
# AstrBot's QQ Official adapter keeps the group/channel scene only
|
||||
# in memory. Restore it before proactive sends so scheduled reports
|
||||
# continue to work after a process restart, before the next event.
|
||||
platform = getattr(self.bot, "platform", None)
|
||||
remember_scene = getattr(platform, "remember_session_scene", None)
|
||||
if callable(remember_scene):
|
||||
remember_scene(str(group_id), "group")
|
||||
umo = f"{self.platform_id}:GroupMessage:{group_id}"
|
||||
return bool(await self._context.send_message(umo, chain))
|
||||
except Exception as exc:
|
||||
logger.error("[QQOfficial] 发送消息失败: %s", exc, exc_info=True)
|
||||
return False
|
||||
|
||||
async def send_text(
|
||||
self, group_id: str, text: str, reply_to: str | None = None
|
||||
) -> bool:
|
||||
from astrbot.api.event import MessageChain
|
||||
|
||||
return await self._send_chain(group_id, MessageChain().message(str(text)))
|
||||
|
||||
async def send_text_report(
|
||||
self,
|
||||
group_id: str,
|
||||
content: str,
|
||||
fallback_content: str | None = None,
|
||||
) -> bool:
|
||||
"""Send long reports as QQ custom Markdown with plain-text fallback."""
|
||||
chunks = self._split_markdown_report(str(content))
|
||||
if not chunks:
|
||||
return True
|
||||
|
||||
markdown_enabled = True
|
||||
sent_markdown_chunks = 0
|
||||
for chunk in chunks:
|
||||
if markdown_enabled:
|
||||
try:
|
||||
if await self._send_markdown_chunk(group_id, chunk):
|
||||
sent_markdown_chunks += 1
|
||||
continue
|
||||
logger.warning(
|
||||
"[QQOfficial] Markdown 接口未返回成功结果,后续改用普通文本"
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"[QQOfficial] Markdown 报告发送失败,后续改用普通文本: %s",
|
||||
exc,
|
||||
)
|
||||
markdown_enabled = False
|
||||
|
||||
if fallback_content and sent_markdown_chunks == 0:
|
||||
for fallback_chunk in self._split_markdown_report(
|
||||
str(fallback_content)
|
||||
):
|
||||
if not await self.send_text(group_id, fallback_chunk):
|
||||
return False
|
||||
return True
|
||||
|
||||
if not await self.send_text(group_id, chunk):
|
||||
return False
|
||||
return True
|
||||
|
||||
async def _send_markdown_chunk(self, group_id: str, content: str) -> bool:
|
||||
api = getattr(self.bot, "api", None)
|
||||
post_group_message = getattr(api, "post_group_message", None)
|
||||
if not callable(post_group_message):
|
||||
return False
|
||||
|
||||
platform = getattr(self.bot, "platform", None)
|
||||
remember_scene = getattr(platform, "remember_session_scene", None)
|
||||
if callable(remember_scene):
|
||||
remember_scene(str(group_id), "group")
|
||||
|
||||
try:
|
||||
from botpy.types.message import MarkdownPayload
|
||||
|
||||
markdown: Any = MarkdownPayload(content=content)
|
||||
except ImportError:
|
||||
# Allows lightweight test environments while botpy is provided by
|
||||
# AstrBot in production.
|
||||
markdown = {"content": content}
|
||||
|
||||
result = await post_group_message( # type: ignore[arg-type]
|
||||
group_openid=str(group_id),
|
||||
msg_type=2,
|
||||
markdown=markdown,
|
||||
msg_seq=self._next_markdown_msg_seq(),
|
||||
)
|
||||
return result is not None
|
||||
|
||||
def _next_markdown_msg_seq(self) -> int:
|
||||
self._markdown_msg_seq = (self._markdown_msg_seq % 10000) + 1
|
||||
return self._markdown_msg_seq
|
||||
|
||||
def _split_markdown_report(self, content: str) -> list[str]:
|
||||
"""Split Markdown on block boundaries without breaking mention tokens."""
|
||||
normalized = str(content or "").strip()
|
||||
if not normalized:
|
||||
return []
|
||||
|
||||
blocks = re.split(r"\n{2,}", normalized)
|
||||
chunks: list[str] = []
|
||||
current = ""
|
||||
|
||||
def append_piece(piece: str) -> None:
|
||||
nonlocal current
|
||||
candidate = f"{current}\n\n{piece}" if current else piece
|
||||
if len(candidate) <= self.MARKDOWN_CHUNK_SIZE:
|
||||
current = candidate
|
||||
return
|
||||
if current:
|
||||
chunks.append(current)
|
||||
current = piece
|
||||
|
||||
for block in blocks:
|
||||
block = block.strip()
|
||||
if not block:
|
||||
continue
|
||||
if len(block) <= self.MARKDOWN_CHUNK_SIZE:
|
||||
append_piece(block)
|
||||
continue
|
||||
|
||||
lines = block.splitlines() or [block]
|
||||
piece = ""
|
||||
for line in lines:
|
||||
candidate = f"{piece}\n{line}" if piece else line
|
||||
if len(candidate) <= self.MARKDOWN_CHUNK_SIZE:
|
||||
piece = candidate
|
||||
continue
|
||||
if piece:
|
||||
append_piece(piece)
|
||||
while len(line) > self.MARKDOWN_CHUNK_SIZE:
|
||||
split_at = self.MARKDOWN_CHUNK_SIZE
|
||||
mention_start = line.rfind("<@", 0, split_at)
|
||||
mention_end = (
|
||||
line.find(">", mention_start) if mention_start >= 0 else -1
|
||||
)
|
||||
if mention_start >= 0 and mention_end >= split_at:
|
||||
split_at = mention_start or self.MARKDOWN_CHUNK_SIZE
|
||||
append_piece(line[:split_at])
|
||||
line = line[split_at:]
|
||||
piece = line
|
||||
if piece:
|
||||
append_piece(piece)
|
||||
|
||||
if current:
|
||||
chunks.append(current)
|
||||
return chunks
|
||||
|
||||
async def send_image(
|
||||
self, group_id: str, image_path: str, caption: str = ""
|
||||
) -> bool:
|
||||
from astrbot.api.event import MessageChain
|
||||
|
||||
chain = MessageChain()
|
||||
if caption:
|
||||
chain.message(caption)
|
||||
if image_path.startswith("base64://"):
|
||||
chain.base64_image(image_path[len("base64://") :])
|
||||
elif image_path.startswith("data:") and "," in image_path:
|
||||
chain.base64_image(image_path.split(",", 1)[1])
|
||||
elif image_path.startswith(("http://", "https://")):
|
||||
chain.url_image(image_path)
|
||||
else:
|
||||
chain.file_image(os.path.abspath(image_path))
|
||||
return await self._send_chain(group_id, chain)
|
||||
|
||||
async def send_file(
|
||||
self, group_id: str, file_path: str, filename: str | None = None
|
||||
) -> bool:
|
||||
# Prefer the public API; fall back to internal for backward compat.
|
||||
# astrbot.core is not part of the stable contract and may change.
|
||||
from astrbot.api.event import MessageChain
|
||||
from astrbot.core.message.components import File
|
||||
|
||||
name = filename or os.path.basename(file_path) or "report"
|
||||
if file_path.startswith(("http://", "https://")):
|
||||
component = File(name=name, url=file_path)
|
||||
else:
|
||||
component = File(name=name, file=os.path.abspath(file_path))
|
||||
return await self._send_chain(group_id, MessageChain([component]))
|
||||
|
||||
async def get_group_info(self, group_id: str) -> UnifiedGroup | None:
|
||||
return UnifiedGroup(
|
||||
group_id=str(group_id),
|
||||
group_name=str(group_id),
|
||||
platform=self.platform_name,
|
||||
)
|
||||
|
||||
async def get_group_list(self) -> list[str]:
|
||||
if self._plugin_instance and hasattr(
|
||||
self._plugin_instance, "get_seen_group_ids"
|
||||
):
|
||||
try:
|
||||
return await self._plugin_instance.get_seen_group_ids(self.platform_id)
|
||||
except Exception as exc:
|
||||
logger.warning("[QQOfficial] 获取已见群列表失败: %s", exc)
|
||||
return []
|
||||
|
||||
async def get_member_list(self, group_id: str) -> list[UnifiedMember]:
|
||||
return []
|
||||
|
||||
async def get_member_info(
|
||||
self, group_id: str, user_id: str
|
||||
) -> UnifiedMember | None:
|
||||
profile = self._member_profiles.get(str(user_id), {})
|
||||
return UnifiedMember(
|
||||
user_id=str(user_id),
|
||||
nickname=profile.get("nickname", ""),
|
||||
avatar_url=await self.get_user_avatar_url(str(user_id)),
|
||||
)
|
||||
|
||||
async def get_user_avatar_url(self, user_id: str, size: int = 100) -> str | None:
|
||||
normalized_user_id = str(user_id or "").strip()
|
||||
if not normalized_user_id:
|
||||
return None
|
||||
profile = self._member_profiles.get(normalized_user_id, {})
|
||||
remembered_avatar_url = profile.get("avatar_url", "")
|
||||
if remembered_avatar_url:
|
||||
return remembered_avatar_url
|
||||
if not self.appid:
|
||||
return None
|
||||
return self.AVATAR_TEMPLATE.format(
|
||||
appid=quote(self.appid, safe=""),
|
||||
member_openid=quote(normalized_user_id, safe=""),
|
||||
)
|
||||
|
||||
async def get_user_avatar_data(self, user_id: str, size: int = 100) -> str | None:
|
||||
avatar_url = await self.get_user_avatar_url(user_id, size)
|
||||
if not avatar_url:
|
||||
return None
|
||||
try:
|
||||
timeout = aiohttp.ClientTimeout(total=15)
|
||||
async with aiohttp.ClientSession(
|
||||
timeout=timeout, trust_env=True
|
||||
) as session:
|
||||
async with session.get(avatar_url) as response:
|
||||
if response.status != 200:
|
||||
return None
|
||||
payload = await response.read()
|
||||
if not payload:
|
||||
return None
|
||||
mime = "image/png" if payload.startswith(b"\x89PNG") else "image/jpeg"
|
||||
return f"data:{mime};base64,{base64.b64encode(payload).decode('utf-8')}"
|
||||
except Exception as exc:
|
||||
logger.debug("[QQOfficial] 下载头像失败: %s", exc)
|
||||
return None
|
||||
|
||||
async def get_group_avatar_url(self, group_id: str, size: int = 100) -> str | None:
|
||||
return None
|
||||
|
||||
async def batch_get_avatar_urls(
|
||||
self, user_ids: list[str], size: int = 100
|
||||
) -> dict[str, str | None]:
|
||||
unique_ids = list(
|
||||
dict.fromkeys(str(user_id) for user_id in user_ids if user_id)
|
||||
)
|
||||
|
||||
async def get_one(user_id: str) -> tuple[str, str | None]:
|
||||
return user_id, await self.get_user_avatar_url(user_id, size)
|
||||
|
||||
return dict(await asyncio.gather(*(get_one(user_id) for user_id in unique_ids)))
|
||||