Files
zhuiguang-ai/scripts/lib/bot-persona-experiment.mjs

681 lines
22 KiB
JavaScript
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// Bot Persona A/B Testing 共享模块
// 用法:
// import { ensureBotVariants, pickAndAssignVariant, recordAssignmentMetrics } from "./lib/bot-persona-experiment.mjs";
//
// // 1) 给 bot 准备至少 2 个变体(首次自动创建 control + variant_a)
// await ensureBotVariants(botConfigId, botUserId, persona);
//
// // 2) 发帖/回复时分流,记录 assignment
// const { variant, assignment } = await pickAndAssignVariant(botConfigId, "topic", topicId, forumSlug);
//
// // 3) 指标采集(在 cron 里调用)
// await recordAssignmentMetrics();
import { PrismaClient } from "@prisma/client";
import { PrismaMariaDb } from "@prisma/adapter-mariadb";
import "dotenv/config";
const base = (process.env.DATABASE_URL || "").replace("mysql://", "mariadb://");
const sep = base.includes("?") ? "&" : "?";
const connectionString = `${base}${sep}connection_limit=3&pool_timeout=10`;
let _prisma = null;
function getPrisma() {
if (!_prisma) {
const adapter = new PrismaMariaDb(connectionString);
_prisma = new PrismaClient({ adapter });
}
return _prisma;
}
function ts() {
return `[${new Date().toISOString()}]`;
}
function pick(arr) {
return arr[Math.floor(Math.random() * arr.length)];
}
function weightedPick(items, weightKey = "weight") {
if (!items || items.length === 0) return null;
const active = items.filter((i) => i.isActive);
const pool = active.length > 0 ? active : items;
const total = pool.reduce((s, x) => s + Math.max(0, x[weightKey] || 0), 0);
if (total <= 0) return pool[0];
let r = Math.random() * total;
for (const it of pool) {
r -= Math.max(0, it[weightKey] || 0);
if (r <= 0) return it;
}
return pool[pool.length - 1];
}
// ======== 变体生成 ========
// 根据现有 persona 派生 2 个变体
// - control: 保留原始行为(无 styleHints)
// - variant_a: "口语化 + 故事化" 风格
// - variant_b: "数据驱动 + 简洁" 风格
// - variant_c(可选): "反主流观点 + 提问引导" 风格
const VARIANT_PRESETS = [
{
variantKey: "control",
label: "对照组(原版画像)",
description: "保留原始 persona 行为,作为 A/B 测试的基准",
isControl: true,
weight: 0.4,
styleHints: {},
},
{
variantKey: "variant_a",
label: "A · 口语化+故事化",
description: "更接地气、加入小故事、问号更多,更像真人聊天",
isControl: false,
weight: 0.3,
styleHints: {
tone: "casual",
questionBoost: 0.2,
storyBoost: 0.3,
maxLength: 400,
forbid: ["综上所述", "从以下几个方面", "首先其次最后", "不可忽视"],
},
},
{
variantKey: "variant_b",
label: "B · 数据驱动+简洁",
description: "多用数据/数字,结构化表达,字数更短",
isControl: false,
weight: 0.3,
styleHints: {
tone: "data_driven",
questionBoost: -0.2,
dataBoost: 0.4,
maxLength: 280,
requireData: true,
forbid: ["综上所述", "从以下几个方面", "首先其次最后"],
},
},
];
/**
* 给一个 bot 创建默认变体(如果还没有)
* - 读现有 BotPersona
* - 给 stanceKeywords、topTopicTypes 做变体化处理
* - 全部 upsert 到 bot_persona_variants 表
* @returns { variants: [...], created: number }
*/
export async function ensureBotVariants(botConfigId, botUserId) {
const prisma = getPrisma();
const existing = await prisma.botPersonaVariant.findMany({
where: { botId: botConfigId },
});
if (existing.length > 0) {
return { variants: existing, created: 0 };
}
const persona = await prisma.botPersona.findUnique({ where: { botId: botConfigId } });
const baseStance = (persona?.stanceKeywords || []).map((k) => k.key).slice(0, 12);
const baseTopicTypes = (persona?.topTopicTypes || []).map((t) => t.name).slice(0, 5);
const created = [];
for (const preset of VARIANT_PRESETS) {
// 不同变体微调关键词:control 全保留;A 加上情绪/故事关键词;B 加上数据关键词
const keywords = [...baseStance];
if (preset.variantKey === "variant_a") {
keywords.push("经历", "故事", "身边", "我", "朋友", "当时", "后来");
} else if (preset.variantKey === "variant_b") {
keywords.push("数据", "比例", "增长", "对比", "案例", "%", "GMV");
}
const v = await prisma.botPersonaVariant.create({
data: {
botId: botConfigId,
variantKey: preset.variantKey,
label: preset.label,
description: preset.description,
isControl: preset.isControl,
isActive: true,
weight: preset.weight,
stanceKeywords: Array.from(new Set(keywords)).slice(0, 15),
topTopicTypes: baseTopicTypes,
topHumanUsers: null,
styleHints: preset.styleHints,
},
});
created.push(v);
}
return { variants: created, created: created.length };
}
/**
* 为某个 bot 启动一个 A/B 实验(默认包含 control + variant_a + variant_b)
* - 多次调用安全:若已存在 status=active 的实验,复用
* - 若已有变体但未挂到实验,挂到该实验
*/
export async function ensureExperiment(botConfigId) {
const prisma = getPrisma();
const existing = await prisma.botPersonaExperiment.findFirst({
where: { botId: botConfigId, status: "active" },
});
if (existing) return existing;
const exp = await prisma.botPersonaExperiment.create({
data: {
botId: botConfigId,
name: "persona A/B",
description: "自动为 bot 启动的 persona A/B 测试:control vs variant_a(口语+故事) vs variant_b(数据+简洁)",
status: "active",
minSamplesPerArm: 20,
significanceLevel: 0.1,
primaryMetric: "engagement_score",
startedAt: new Date(),
},
});
// 把该 bot 已有变体挂到实验
await prisma.botPersonaVariant.updateMany({
where: { botId: botConfigId, experimentId: null },
data: { experimentId: exp.id },
});
return exp;
}
/**
* 给某条内容(topic / post)挑选一个变体并记录 assignment
* @param botConfigId
* @param refType "topic" | "post"
* @param refId topicId / postId(必须 > 0;先选变体再创建内容时使用 pickVariantForPrompt)
* @param forumSlug 可选
* @returns { variant, assignment, personaBlock }
*/
export async function pickAndAssignVariant(botConfigId, refType, refId, forumSlug = null) {
const prisma = getPrisma();
// 1) 确保变体存在
const { variants } = await ensureBotVariants(botConfigId);
if (variants.length === 0) {
return { variant: null, assignment: null, personaBlock: "" };
}
// 2) 选变体(按 weight)
const variant = weightedPick(variants, "weight");
if (!variant) {
return { variant: null, assignment: null, personaBlock: "" };
}
// 3) 写 assignment(unique key 保护)
const assignment = await prisma.botPersonaAssignment.upsert({
where: { refType_refId: { refType, refId } },
create: {
botId: botConfigId,
variantId: variant.id,
refType,
refId,
forumSlug,
},
update: {
// 同一 ref 已存在,不覆盖 variant(保持首次分流)
botId: botConfigId,
},
});
// 4) 变体 sample + 1
await prisma.botPersonaVariant.update({
where: { id: variant.id },
data: { sampleCount: { increment: 1 } },
});
return {
variant,
assignment,
personaBlock: buildVariantPersonaBlock(variant),
};
}
/**
* 轻量版:只选变体 + 返回 personaBlock,**不写库**
* 适用于"先生成内容再创建记录的顺序"场景(如先调用 LLM 再 createTopic)
* 配合 commitAssignment(botConfigId, refType, refId, variantId) 在内容创建后落盘
*/
export async function pickVariantForPrompt(botConfigId) {
const prisma = getPrisma();
const { variants } = await ensureBotVariants(botConfigId);
if (variants.length === 0) {
return { variant: null, personaBlock: "" };
}
const variant = weightedPick(variants, "weight");
if (!variant) {
return { variant: null, personaBlock: "" };
}
return {
variant,
personaBlock: buildVariantPersonaBlock(variant),
};
}
/**
* 提交 assignment(轻量版的配套写入)
* - 幂等:同一 (refType, refId) 多次调用不会变 variant
* - 会自增变体 sampleCount
*/
export async function commitAssignment(botConfigId, variantId, refType, refId, forumSlug = null) {
if (!variantId || !refId) return null;
const prisma = getPrisma();
const assignment = await prisma.botPersonaAssignment.upsert({
where: { refType_refId: { refType, refId } },
create: {
botId: botConfigId,
variantId,
refType,
refId,
forumSlug,
},
update: { botId: botConfigId },
});
await prisma.botPersonaVariant.update({
where: { id: variantId },
data: { sampleCount: { increment: 1 } },
});
return assignment;
}
// 把变体的 styleHints 格式化成可注入 prompt 的 block
export function buildVariantPersonaBlock(variant) {
if (!variant) return "";
const hints = variant.styleHints || {};
const parts = [];
parts.push(`[A/B 变体 ${variant.variantKey} · ${variant.label}]`);
if (hints.tone) {
const toneMap = {
casual: "语气:口语化、接地气",
data_driven: "语气:数据驱动、引用具体数字",
contrarian: "语气:反主流观点、犀利但有理",
};
parts.push(toneMap[hints.tone] || `语气:${hints.tone}`);
}
if (typeof hints.maxLength === "number") {
parts.push(`字数上限:约 ${hints.maxLength} 字`);
}
if (hints.questionBoost) {
parts.push(`问号倾向:${hints.questionBoost > 0 ? "略多问号" : "少用问句"}`);
}
if (hints.dataBoost) parts.push("多用数据/数字");
if (hints.storyBoost) parts.push("多讲小故事/经历");
if (Array.isArray(hints.forbid) && hints.forbid.length > 0) {
parts.push(`避免:${hints.forbid.join("、")}`);
}
if (variant.stanceKeywords && variant.stanceKeywords.length > 0) {
parts.push(`关注关键词:${variant.stanceKeywords.slice(0, 8).join("、")}`);
}
return `\n${parts.join("\n")}\n`;
}
// ======== 指标采集 ========
/**
* 把每条 assignment 的实际互动(replyCount/likeCount/humanReplies)回填,
* 并按 (variantId, date) 写入 bot_persona_metrics 聚合
* - 默认只看过去 7 天的 assignment(性能保护)
*/
export async function recordAssignmentMetrics(options = {}) {
const prisma = getPrisma();
const lookbackDays = options.lookbackDays ?? 7;
const since = new Date(Date.now() - lookbackDays * 24 * 60 * 60 * 1000);
// 1) 拉所有变体(活跃实验)
const variants = await prisma.botPersonaVariant.findMany({
where: { isActive: true },
select: { id: true, botId: true },
});
if (variants.length === 0) return { updated: 0 };
const variantIds = variants.map((v) => v.id);
const assignments = await prisma.botPersonaAssignment.findMany({
where: {
variantId: { in: variantIds },
createdAt: { gte: since },
},
});
if (assignments.length === 0) return { updated: 0 };
// 2) 按 refType 拆开批量查真实互动数
const topicIds = assignments.filter((a) => a.refType === "topic").map((a) => a.refId);
const postIds = assignments.filter((a) => a.refType === "post").map((a) => a.refId);
// topic: 真实 replyCount / likeCount / 真人 reply 数
const topicStats = await prisma.forumTopic.findMany({
where: { id: { in: topicIds.length > 0 ? topicIds : [-1] } },
select: {
id: true,
replyCount: true,
likeCount: true,
viewCount: true,
posts: {
where: { user: { isBot: false } },
select: { id: true },
},
},
});
const topicMap = new Map(
topicStats.map((t) => [
t.id,
{
replies: t.replyCount || 0,
likes: t.likeCount || 0,
impressions: t.viewCount || 0,
humanReplies: t.posts.length,
},
])
);
// post: 真实 likeCount
const postStats = await prisma.forumPost.findMany({
where: { id: { in: postIds.length > 0 ? postIds : [-1] } },
select: { id: true, likeCount: true, topic: { select: { id: true } } },
});
const postMap = new Map(postStats.map((p) => [p.id, p.likeCount || 0]));
// 对 post 型 assignment,人工统计"该 post 之后同 topic 的新回复":用 topicId+postCreatedAt 之后的人类回复数
// 简化:取 topic 的总 replyCount / totalPosts 数 - 该 post 之前的 - 1 = 该 post 之后的新回复
// 为简化与一致性:post 的 replies/humanReplies 用同 topic 下的总数估算,impressions=0
// (post 的真实互动其实主要看 likeCount)
// 3) 写回 assignment
let updated = 0;
for (const a of assignments) {
let replies = 0;
let likes = 0;
let impressions = 0;
let humanReplies = 0;
if (a.refType === "topic" && topicMap.has(a.refId)) {
const t = topicMap.get(a.refId);
replies = t.replies;
likes = t.likes;
impressions = t.impressions;
humanReplies = t.humanReplies;
} else if (a.refType === "post") {
likes = postMap.get(a.refId) || 0;
// post 的回复/曝光粗略处理:0
replies = 0;
humanReplies = 0;
impressions = 0;
}
if (
replies !== a.replies ||
likes !== a.likes ||
humanReplies !== a.humanReplies ||
impressions !== a.impressions
) {
await prisma.botPersonaAssignment.update({
where: { id: a.id },
data: { replies, likes, humanReplies, impressions },
});
updated++;
}
}
// 4) 写入按 (variantId, date) 聚合的 metrics
const dayBuckets = new Map(); // key: `${variantId}|${yyyy-mm-dd}` -> { variantId, date, contentCount, replies, likes, humanReplies, impressions }
for (const a of assignments) {
const dateStr = a.createdAt.toISOString().slice(0, 10);
const key = `${a.variantId}|${dateStr}`;
if (!dayBuckets.has(key)) {
dayBuckets.set(key, {
variantId: a.variantId,
date: new Date(`${dateStr}T00:00:00.000Z`),
contentCount: 0,
replies: 0,
likes: 0,
humanReplies: 0,
impressions: 0,
});
}
const b = dayBuckets.get(key);
b.contentCount += 1;
b.replies += a.replies;
b.likes += a.likes;
b.humanReplies += a.humanReplies;
b.impressions += a.impressions;
}
for (const b of dayBuckets.values()) {
const avgReplies = b.contentCount > 0 ? b.replies / b.contentCount : 0;
const avgLikes = b.contentCount > 0 ? b.likes / b.contentCount : 0;
// 综合得分:平均点赞×3 + 真人回复×5 + 平均回复×1
const computedScore = avgLikes * 3 + b.humanReplies / Math.max(1, b.contentCount) * 5 + avgReplies * 1;
await prisma.botPersonaMetric.upsert({
where: { variantId_date: { variantId: b.variantId, date: b.date } },
create: {
variantId: b.variantId,
date: b.date,
impressions: b.impressions,
replies: b.replies,
likes: b.likes,
humanReplies: b.humanReplies,
contentCount: b.contentCount,
avgReplies,
avgLikes,
computedScore,
},
update: {
impressions: b.impressions,
replies: b.replies,
likes: b.likes,
humanReplies: b.humanReplies,
contentCount: b.contentCount,
avgReplies,
avgLikes,
computedScore,
},
});
}
// 5) 回写变体的累计统计(最近 7 天 assignment 滚动聚合)
const variantAgg = new Map();
for (const a of assignments) {
if (!variantAgg.has(a.variantId)) {
variantAgg.set(a.variantId, { sampleCount: 0, replyCount: 0, likeCount: 0, humanReplyCount: 0 });
}
const v = variantAgg.get(a.variantId);
v.sampleCount += 1;
v.replyCount += a.replies;
v.likeCount += a.likes;
v.humanReplyCount += a.humanReplies;
}
for (const [variantId, agg] of variantAgg.entries()) {
// 综合分:平均每篇互动 = (收赞 + 收真人回复 + 收回复) / 样本
const score = agg.sampleCount > 0
? (agg.likeCount * 3 + agg.humanReplyCount * 5 + agg.replyCount) / agg.sampleCount
: 0;
await prisma.botPersonaVariant.update({
where: { id: variantId },
data: {
sampleCount: agg.sampleCount,
replyCount: agg.replyCount,
likeCount: agg.likeCount,
humanReplyCount: agg.humanReplyCount,
engagementScore: Math.round(score * 100) / 100,
},
});
}
return { updated, assignments: assignments.length, dayBuckets: dayBuckets.size };
}
// ======== 分析 / 变体切换 ========
/**
* 对单个 bot 的所有变体做显著性分析
* - 输入:变体列表(含 sampleCount/replyCount/likeCount/humanReplyCount/engagementScore)
* - 输出:winnerVariantKey / analysis 详情
* 策略:双比例 z 检验(高斯近似),比较每个 variant vs control 的互动率((replies+likes+humanReplies)/sampleCount)
*/
export function analyzeBotVariants(variants) {
if (!variants || variants.length === 0) {
return { winner: null, reason: "no_variants", comparisons: [] };
}
const control = variants.find((v) => v.isControl) || variants[0];
const others = variants.filter((v) => v.id !== control.id);
if (control.sampleCount < 1) {
return { winner: null, reason: "control_no_samples", comparisons: [] };
}
const controlRate = (control.likeCount * 3 + control.humanReplyCount * 5 + control.replyCount) / control.sampleCount;
const controlN = control.sampleCount;
const comparisons = others.map((v) => {
if (v.sampleCount < 1) {
return { variant: v, rate: 0, pValue: null, significant: false, winner: false };
}
const rate = (v.likeCount * 3 + v.humanReplyCount * 5 + v.replyCount) / v.sampleCount;
const n = v.sampleCount;
// pooled proportion
const pooled = (control.likeCount * 3 + control.humanReplyCount * 5 + control.replyCount + v.likeCount * 3 + v.humanReplyCount * 5 + v.replyCount) / (controlN + n);
const se = Math.sqrt(pooled * (1 - pooled) * (1 / controlN + 1 / n));
let z = 0;
let p = 1;
if (se > 0) {
z = (rate - controlRate) / se;
// 双侧检验
p = 2 * (1 - normalCdf(Math.abs(z)));
}
return {
variant: v,
rate: Math.round(rate * 1000) / 1000,
controlRate: Math.round(controlRate * 1000) / 1000,
zScore: Math.round(z * 100) / 100,
pValue: Math.round(p * 1000) / 1000,
significant: p < 0.1 && n >= 20, // significanceLevel=0.1, minSamples=20
winner: false,
};
});
// 选 winner:所有"显著更优"的里面 p 最小 + sample 最多的
const winners = comparisons.filter((c) => c.significant && c.rate > c.controlRate);
let winner = null;
if (winners.length > 0) {
winners.sort((a, b) => a.pValue - b.pValue);
winners[0].winner = true;
winner = winners[0].variant;
}
return {
winner,
winnerReason: winner
? `variant ${winner.variantKey} 综合互动率显著高于 control (p=${winners[0].pValue})`
: "no significant winner yet",
control,
comparisons,
};
}
// 标准正态分布 CDF(Abramowitz & Stegun 近似)
function normalCdf(x) {
const t = 1 / (1 + 0.2316419 * x);
const d = 0.3989422804014327 * Math.exp(-x * x / 2);
let p = d * t * (0.319381530 + t * (-0.356563782 + t * (1.781477937 + t * (-1.821255978 + t * 1.330274429))));
return 1 - p;
}
/**
* 完整跑一遍分析 + 切换
* - 遍历所有 bot
* - 对每个 bot 调 analyzeBotVariants
* - 若有 winner 显著更优,把实验标记为 completed
* - 后续 bot 在选变体时仍会按 weight 选(不会自动锁死一个变体),但 recordAssignmentMetrics 会持续跑
* - 还可以加:winner 出现后,把 winner 权重提到 1.0,其他降到 0
*/
export async function analyzeAndSwitchAllBots(options = {}) {
const prisma = getPrisma();
const onlyActive = options.onlyActive ?? true;
// 取所有有变体的 bot
const variants = await prisma.botPersonaVariant.findMany({
where: onlyActive ? { isActive: true } : undefined,
include: { experiment: true },
});
const byBot = new Map();
for (const v of variants) {
if (!byBot.has(v.botId)) byBot.set(v.botId, []);
byBot.get(v.botId).push(v);
}
const results = [];
for (const [botId, list] of byBot.entries()) {
const analysis = analyzeBotVariants(list);
if (analysis.winner) {
const exp = list.find((v) => v.experimentId)?.experiment;
if (exp && exp.status === "active") {
// 标记实验完成
await prisma.botPersonaExperiment.update({
where: { id: exp.id },
data: {
status: "completed",
winnerVariantKey: analysis.winner.variantKey,
endedAt: new Date(),
lastAnalyzedAt: new Date(),
},
});
// 提升 winner 权重到 0.85,其他活跃变体降到 0.075
await prisma.botPersonaVariant.update({
where: { id: analysis.winner.id },
data: { weight: 0.85 },
});
await prisma.botPersonaVariant.updateMany({
where: {
botId,
id: { not: analysis.winner.id },
experimentId: exp.id,
},
data: { weight: 0.075 },
});
results.push({
botId,
experimentId: exp.id,
winner: analysis.winner.variantKey,
reason: analysis.winnerReason,
action: "switched_weights",
});
} else {
results.push({
botId,
winner: analysis.winner.variantKey,
reason: analysis.winnerReason,
action: "no_active_experiment",
});
}
} else {
// 没有 winner,刷新 lastAnalyzedAt
const exp = list.find((v) => v.experimentId)?.experiment;
if (exp && exp.status === "active") {
await prisma.botPersonaExperiment.update({
where: { id: exp.id },
data: { lastAnalyzedAt: new Date() },
});
}
results.push({
botId,
winner: null,
reason: analysis.winnerReason,
action: "monitoring",
});
}
}
return { botCount: byBot.size, results };
}
/**
* 优雅关闭
*/
export async function disconnectBotExperiment() {
if (_prisma) {
await _prisma.$disconnect();
_prisma = null;
}
}