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>
This commit is contained in:
2026-09-01 13:13:40 +08:00
co-authored by Claude
parent 1783c60afa
commit b61d09f09f
3201 changed files with 160436 additions and 171 deletions
@@ -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()
File diff suppressed because one or more lines are too long
@@ -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
Binary file not shown.

After

Width:  |  Height:  |  Size: 540 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 726 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 690 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 1.7 MiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 913 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 793 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 951 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 1.1 MiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 630 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 1.2 MiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 1.3 MiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 271 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 677 KiB

Binary file not shown.

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"
}
]
}
Binary file not shown.

After

Width:  |  Height:  |  Size: 843 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 1.0 MiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 629 KiB

Binary file not shown.

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)))

Some files were not shown because too many files have changed in this diff Show More