feat: 提示词和 AI 参数移至数据库管理

- AiSetting 模型新增 6 个提示词字段 + 3 个阶段温度字段
- main.py 迁移逻辑改为检测 ai_settings 表的新字段并填充默认值
- generate.py 从数据库读取生成提示词,去掉硬编码
- optimizer 从 API 读取各阶段提示词和温度,删除 prompts.py
- crawler 从 API 读取提取/改写提示词和温度,删除 prompts.py
- settings/active 端点去掉 token 认证(供爬虫/优化器使用)
- 后台设置页新增提示词编辑区和温度调节控件
- 新增 _ensure_default_settings 自动创建默认配置
This commit is contained in:
bwstudio
2026-06-13 17:27:09 +08:00
parent 6d015bb56d
commit db8cea02f8
11 changed files with 457 additions and 195 deletions
+105 -3
View File
@@ -2,10 +2,10 @@
<div class="setting-page"> <div class="setting-page">
<h2>AI 配置</h2> <h2>AI 配置</h2>
<!-- AI 配置表单 --> <!-- AI 基础配置 -->
<el-card style="margin-top: 20px"> <el-card style="margin-top: 20px">
<template #header> <template #header>
<span>AI 配置</span> <span>API 基础配置</span>
</template> </template>
<el-form :model="form" label-width="120px"> <el-form :model="form" label-width="120px">
<el-form-item label="提供商"> <el-form-item label="提供商">
@@ -57,6 +57,81 @@
</el-form> </el-form>
</el-card> </el-card>
<!-- 提示词配置 -->
<el-card style="margin-top: 20px">
<template #header>
<span>提示词模板</span>
</template>
<el-form :model="form" label-width="140px">
<el-form-item label="AI 生成笑话">
<el-input
v-model="form.generate_prompt"
type="textarea"
:rows="6"
placeholder="使用 {context}/{style}/{length_requirement} 作为占位符"
style="width: 100%"
/>
</el-form-item>
<el-divider />
<el-form-item label="优化器-质量检测">
<el-input
v-model="form.optimizer_quality_prompt"
type="textarea"
:rows="4"
style="width: 100%"
/>
</el-form-item>
<el-form-item label="质量检测温度">
<el-input-number v-model="form.optimizer_quality_temperature" :min="0" :max="2" :step="0.1" />
</el-form-item>
<el-divider />
<el-form-item label="优化器-润色">
<el-input
v-model="form.optimizer_polish_prompt"
type="textarea"
:rows="4"
style="width: 100%"
/>
</el-form-item>
<el-form-item label="润色温度">
<el-input-number v-model="form.optimizer_polish_temperature" :min="0" :max="2" :step="0.1" />
</el-form-item>
<el-divider />
<el-form-item label="优化器-评价分类">
<el-input
v-model="form.optimizer_evaluate_prompt"
type="textarea"
:rows="4"
style="width: 100%"
/>
</el-form-item>
<el-form-item label="评价温度">
<el-input-number v-model="form.optimizer_evaluate_temperature" :min="0" :max="2" :step="0.1" />
</el-form-item>
<el-divider />
<el-form-item label="爬虫-提取笑话">
<el-input
v-model="form.crawler_extract_prompt"
type="textarea"
:rows="4"
style="width: 100%"
/>
</el-form-item>
<el-form-item label="爬虫-改写">
<el-input
v-model="form.crawler_rewrite_prompt"
type="textarea"
:rows="4"
style="width: 100%"
/>
</el-form-item>
</el-form>
</el-card>
<div style="margin-top: 20px"> <div style="margin-top: 20px">
<el-button type="primary" @click="handleSave" :loading="saving">保存配置</el-button> <el-button type="primary" @click="handleSave" :loading="saving">保存配置</el-button>
</div> </div>
@@ -106,6 +181,15 @@ const form = ref({
crawl_enabled: false, crawl_enabled: false,
crawl_keywords: '冷笑话,段子,谐音梗', crawl_keywords: '冷笑话,段子,谐音梗',
max_pages_per_run: 3, max_pages_per_run: 3,
generate_prompt: '',
optimizer_quality_prompt: '',
optimizer_polish_prompt: '',
optimizer_evaluate_prompt: '',
optimizer_quality_temperature: 0.3,
optimizer_polish_temperature: 0.8,
optimizer_evaluate_temperature: 0.3,
crawler_extract_prompt: '',
crawler_rewrite_prompt: '',
}) })
const allSettings = ref([]) const allSettings = ref([])
@@ -126,6 +210,15 @@ const loadData = async () => {
crawl_enabled: active.crawl_enabled, crawl_enabled: active.crawl_enabled,
crawl_keywords: active.crawl_keywords || '', crawl_keywords: active.crawl_keywords || '',
max_pages_per_run: active.max_pages_per_run, max_pages_per_run: active.max_pages_per_run,
generate_prompt: active.generate_prompt || '',
optimizer_quality_prompt: active.optimizer_quality_prompt || '',
optimizer_polish_prompt: active.optimizer_polish_prompt || '',
optimizer_evaluate_prompt: active.optimizer_evaluate_prompt || '',
optimizer_quality_temperature: active.optimizer_quality_temperature ?? 0.3,
optimizer_polish_temperature: active.optimizer_polish_temperature ?? 0.8,
optimizer_evaluate_temperature: active.optimizer_evaluate_temperature ?? 0.3,
crawler_extract_prompt: active.crawler_extract_prompt || '',
crawler_rewrite_prompt: active.crawler_rewrite_prompt || '',
} }
} catch (e) { } catch (e) {
// 没有激活配置,使用默认值 // 没有激活配置,使用默认值
@@ -147,7 +240,7 @@ const handleSave = async () => {
} else { } else {
const created = await createSetting(form.value) const created = await createSetting(form.value)
editingId.value = created.id editingId.value = created.id
await toggleSetting(created.id) // 设为激活 await toggleSetting(created.id)
ElMessage.success('创建并激活成功') ElMessage.success('创建并激活成功')
} }
await loadData() await loadData()
@@ -170,6 +263,15 @@ const handleEdit = (row) => {
crawl_enabled: row.crawl_enabled, crawl_enabled: row.crawl_enabled,
crawl_keywords: row.crawl_keywords || '', crawl_keywords: row.crawl_keywords || '',
max_pages_per_run: row.max_pages_per_run, max_pages_per_run: row.max_pages_per_run,
generate_prompt: row.generate_prompt || '',
optimizer_quality_prompt: row.optimizer_quality_prompt || '',
optimizer_polish_prompt: row.optimizer_polish_prompt || '',
optimizer_evaluate_prompt: row.optimizer_evaluate_prompt || '',
optimizer_quality_temperature: row.optimizer_quality_temperature ?? 0.3,
optimizer_polish_temperature: row.optimizer_polish_temperature ?? 0.8,
optimizer_evaluate_temperature: row.optimizer_evaluate_temperature ?? 0.3,
crawler_extract_prompt: row.crawler_extract_prompt || '',
crawler_rewrite_prompt: row.crawler_rewrite_prompt || '',
} }
} }
+13
View File
@@ -14,6 +14,19 @@ class AiSetting(Base):
temperature = Column(Float, nullable=False, default=0.7) temperature = Column(Float, nullable=False, default=0.7)
max_tokens = Column(Integer, nullable=False, default=2048) max_tokens = Column(Integer, nullable=False, default=2048)
# 提示词模板
generate_prompt = Column(Text, nullable=True, comment="AI 笑话生成提示词")
optimizer_quality_prompt = Column(Text, nullable=True, comment="优化器-质量检测提示词")
optimizer_polish_prompt = Column(Text, nullable=True, comment="优化器-润色提示词")
optimizer_evaluate_prompt = Column(Text, nullable=True, comment="优化器-评价分类提示词")
crawler_extract_prompt = Column(Text, nullable=True, comment="爬虫-笑话提取提示词")
crawler_rewrite_prompt = Column(Text, nullable=True, comment="爬虫-改写提示词")
# 优化器各阶段温度
optimizer_quality_temperature = Column(Float, nullable=False, default=0.3)
optimizer_polish_temperature = Column(Float, nullable=False, default=0.8)
optimizer_evaluate_temperature = Column(Float, nullable=False, default=0.3)
crawl_enabled = Column(Boolean, nullable=False, default=False) crawl_enabled = Column(Boolean, nullable=False, default=False)
crawl_keywords = Column(Text, nullable=True, default="") crawl_keywords = Column(Text, nullable=True, default="")
max_pages_per_run = Column(Integer, nullable=False, default=3) max_pages_per_run = Column(Integer, nullable=False, default=3)
+33 -33
View File
@@ -1,8 +1,8 @@
"""智能笑话生成器 API""" """AI 笑话生成器 API — 提示词从数据库读取"""
import json import json
from datetime import datetime from datetime import datetime
from fastapi import APIRouter, Depends, HTTPException, Query from fastapi import APIRouter, Depends, HTTPException
from openai import OpenAI from openai import OpenAI
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
@@ -13,9 +13,31 @@ from app.schemas.joke import GenerateRequest, GenerateResponse
router = APIRouter(prefix="/generate", tags=["生成器"]) router = APIRouter(prefix="/generate", tags=["生成器"])
STYLE_MAP = {
"cold": "冷幽默 / 无厘头",
"warm": "温馨幽默 / 暖心搞笑",
"twist": "反转 / 神转折",
"pun": "谐音梗 / 文字游戏",
"sketch": "段子 / 吐槽调侃",
"irony": "讽刺幽默 / 黑色幽默",
}
# AI Prompt LENGTH_MAP = {
GENERATION_PROMPT = """你是一位幽默大师,专门创作轻松搞笑的短笑话。 "short": "30-80字,非常简短",
"medium": "80-150字,正常长度",
"long": "150-300字,可以描述一个小场景",
}
def _load_prompt(setting: AiSetting) -> str:
"""从数据库读取生成提示词,没有则返回默认值"""
prompt = setting.generate_prompt
if not prompt or not prompt.strip():
prompt = GENERATE_PROMPT_DEFAULT
return prompt
GENERATE_PROMPT_DEFAULT = """你是一位幽默大师,专门创作轻松搞笑的短笑话。
{context} {context}
@@ -36,27 +58,13 @@ GENERATION_PROMPT = """你是一位幽默大师,专门创作轻松搞笑的短
score 是 1-10 的整数评分,reason 是用一句话说明亮点。""" score 是 1-10 的整数评分,reason 是用一句话说明亮点。"""
STYLE_MAP = {
"cold": "冷幽默 / 无厘头",
"warm": "温馨幽默 / 暖心搞笑",
"twist": "反转 / 神转折",
"pun": "谐音梗 / 文字游戏",
"sketch": "段子 / 吐槽调侃",
"irony": "讽刺幽默 / 黑色幽默",
}
LENGTH_MAP = {
"short": "30-80字,非常简短",
"medium": "80-150字,正常长度",
"long": "150-300字,可以描述一个小场景",
}
def _build_prompt( def _build_prompt(
scenarios: list[str], scenarios: list[str],
keywords: list[str], keywords: list[str],
style: str = "twist", style: str = "twist",
length: str = "medium", length: str = "medium",
setting: AiSetting | None = None,
) -> str: ) -> str:
"""构建 AI prompt,支持风格和长度控制""" """构建 AI prompt,支持风格和长度控制"""
parts = [] parts = []
@@ -66,7 +74,9 @@ def _build_prompt(
parts.append(f"关键词:{', '.join(keywords)}") parts.append(f"关键词:{', '.join(keywords)}")
if not parts: if not parts:
parts.append("场景:日常生活的各种趣事(不指定具体场景)") parts.append("场景:日常生活的各种趣事(不指定具体场景)")
return GENERATION_PROMPT.format(
prompt_template = _load_prompt(setting) if setting else GENERATE_PROMPT_DEFAULT
return prompt_template.format(
context="\n".join(parts), context="\n".join(parts),
style=STYLE_MAP.get(style, "反转 / 神转折"), style=STYLE_MAP.get(style, "反转 / 神转折"),
length_requirement=LENGTH_MAP.get(length, "80-150字,正常长度"), length_requirement=LENGTH_MAP.get(length, "80-150字,正常长度"),
@@ -78,7 +88,7 @@ def generate_joke(
req: GenerateRequest, req: GenerateRequest,
db: Session = Depends(get_db), db: Session = Depends(get_db),
): ):
"""调用 AI 生成笑话,支持风格和长度控制""" """调用 AI 生成笑话,提示词从数据库读取"""
# 获取激活的 AI 配置 # 获取激活的 AI 配置
setting = db.query(AiSetting).filter(AiSetting.is_active == True).first() setting = db.query(AiSetting).filter(AiSetting.is_active == True).first()
if not setting: if not setting:
@@ -87,10 +97,9 @@ def generate_joke(
if not setting or not setting.api_key: if not setting or not setting.api_key:
raise HTTPException(status_code=503, detail="AI 服务未配置,请联系管理员") raise HTTPException(status_code=503, detail="AI 服务未配置,请联系管理员")
# 调用 AI
try: try:
client = OpenAI(base_url=setting.api_base, api_key=setting.api_key) client = OpenAI(base_url=setting.api_base, api_key=setting.api_key)
prompt = _build_prompt(req.scenarios, req.keywords, req.style, req.length) prompt = _build_prompt(req.scenarios, req.keywords, req.style, req.length, setting)
response = client.chat.completions.create( response = client.chat.completions.create(
model=setting.model_name, model=setting.model_name,
@@ -113,22 +122,18 @@ def _parse_json_response(raw: str | None) -> GenerateResponse:
if not raw or not raw.strip(): if not raw or not raw.strip():
return GenerateResponse(title="生成的笑话", content="(内容生成失败,请重新生成)", score=0) return GenerateResponse(title="生成的笑话", content="(内容生成失败,请重新生成)", score=0)
# 尝试从 markdown 代码块中提取 JSON
raw = raw.strip() raw = raw.strip()
if raw.startswith("```"): if raw.startswith("```"):
lines = raw.split("\n") lines = raw.split("\n")
# 去掉第一行 ```json 和最后一行 ```
if len(lines) >= 3: if len(lines) >= 3:
raw = "\n".join(lines[1:-1]).strip() raw = "\n".join(lines[1:-1]).strip()
# 移除可能的尾部分号
if raw.endswith(","): if raw.endswith(","):
raw = raw[:-1] raw = raw[:-1]
try: try:
data = json.loads(raw) data = json.loads(raw)
except json.JSONDecodeError: except json.JSONDecodeError:
# 如果 JSON 解析失败,尝试查找花括号内的内容
try: try:
start = raw.index("{") start = raw.index("{")
end = raw.rindex("}") + 1 end = raw.rindex("}") + 1
@@ -147,20 +152,17 @@ def _parse_json_response(raw: str | None) -> GenerateResponse:
def _parse_response(raw: str | None) -> GenerateResponse: def _parse_response(raw: str | None) -> GenerateResponse:
"""解析 AI 返回内容,提取标题和内容""" """解析 AI 返回内容(非 JSON 回退)"""
# 防御:处理空或 None 输入
if not raw or not raw.strip(): if not raw or not raw.strip():
raise ValueError("AI 返回内容为空") raise ValueError("AI 返回内容为空")
title = "" title = ""
content = raw content = raw
# 尝试提取 "标题:xxx" 或 "标题:xxx"
for line in raw.split("\n"): for line in raw.split("\n"):
line = line.strip() line = line.strip()
if line.startswith("标题:") or line.startswith("标题:"): if line.startswith("标题:") or line.startswith("标题:"):
title = line.split("", 1)[-1].split(":", 1)[-1].strip() title = line.split("", 1)[-1].split(":", 1)[-1].strip()
# 只替换这一行,不要 replace 全局
lines = content.split("\n") lines = content.split("\n")
for i, l in enumerate(lines): for i, l in enumerate(lines):
if l.strip() == line: if l.strip() == line:
@@ -169,14 +171,12 @@ def _parse_response(raw: str | None) -> GenerateResponse:
content = "\n".join(lines).strip() content = "\n".join(lines).strip()
break break
# 如果没有提取到标题,取第一行
if not title: if not title:
first_line = raw.split("\n")[0].strip() first_line = raw.split("\n")[0].strip()
if first_line.startswith("标题"): if first_line.startswith("标题"):
first_line = first_line.split("", 1)[-1].split(":", 1)[-1].strip() first_line = first_line.split("", 1)[-1].split(":", 1)[-1].strip()
title = first_line[:30] if len(first_line) > 30 else first_line title = first_line[:30] if len(first_line) > 30 else first_line
# 防御:content 不能为空
if not content.strip(): if not content.strip():
content = "(内容生成失败,请重新生成)" content = "(内容生成失败,请重新生成)"
+22 -4
View File
@@ -22,12 +22,12 @@ def list_settings(
@router.get("/active") @router.get("/active")
def get_active_setting( def get_active_setting(
db: Session = Depends(get_db), db: Session = Depends(get_db),
current_user: AdminUser = Depends(get_current_admin_user),
): ):
"""获取当前激活的 AI 配置(爬虫调用,无需用户认证,token 校验仍保留""" """获取当前激活的 AI 配置(无需认证,供爬虫/优化器使用"""
setting = db.query(AiSetting).filter(AiSetting.is_active == True).first() setting = db.query(AiSetting).filter(AiSetting.is_active == True).first()
if not setting: if not setting:
raise HTTPException(status_code=404, detail="未找到激活的 AI 配置") from app.routers.settings import _ensure_default_settings
setting = _ensure_default_settings(db)
return setting return setting
@@ -93,4 +93,22 @@ def delete_setting(
raise HTTPException(status_code=404, detail="配置不存在") raise HTTPException(status_code=404, detail="配置不存在")
db.delete(db_setting) db.delete(db_setting)
db.commit() db.commit()
return {"message": "删除成功"} return {"message": "删除成功"}
def _ensure_default_settings(db: Session) -> AiSetting:
"""当没有激活配置时,创建一条默认配置"""
existing = db.query(AiSetting).first()
if existing:
existing.is_active = True
db.commit()
return existing
defaults = AiSetting(
is_active=True,
# 默认提示词会在 main.py 迁移时填充
)
db.add(defaults)
db.commit()
db.refresh(defaults)
return defaults
+15 -1
View File
@@ -1,4 +1,4 @@
from pydantic import BaseModel, ConfigDict, Field from pydantic import BaseModel, Field
class AiSettingBase(BaseModel): class AiSettingBase(BaseModel):
@@ -8,6 +8,20 @@ class AiSettingBase(BaseModel):
model_name: str = Field(default="nvidia/llama-3.1-nemotron-70b-instruct") model_name: str = Field(default="nvidia/llama-3.1-nemotron-70b-instruct")
temperature: float = Field(default=0.7, ge=0, le=2) temperature: float = Field(default=0.7, ge=0, le=2)
max_tokens: int = Field(default=2048, ge=1) max_tokens: int = Field(default=2048, ge=1)
# 提示词
generate_prompt: str = Field(default="")
optimizer_quality_prompt: str = Field(default="")
optimizer_polish_prompt: str = Field(default="")
optimizer_evaluate_prompt: str = Field(default="")
crawler_extract_prompt: str = Field(default="")
crawler_rewrite_prompt: str = Field(default="")
# 优化器各阶段温度
optimizer_quality_temperature: float = Field(default=0.3, ge=0, le=2)
optimizer_polish_temperature: float = Field(default=0.8, ge=0, le=2)
optimizer_evaluate_temperature: float = Field(default=0.3, ge=0, le=2)
crawl_enabled: bool = Field(default=False) crawl_enabled: bool = Field(default=False)
crawl_keywords: str = Field(default="") crawl_keywords: str = Field(default="")
max_pages_per_run: int = Field(default=3, ge=1) max_pages_per_run: int = Field(default=3, ge=1)
+108 -5
View File
@@ -13,14 +13,117 @@ from app.models.feedback import Feedback # 注册模型,确保 create_all 能
Base.metadata.create_all(bind=engine) Base.metadata.create_all(bind=engine)
# 对已有表新增字段的兼容迁移(SQLite 不支持 ALTER TABLE ADD COLUMN IF NOT EXISTS # 对已有表新增字段的兼容迁移(SQLite 不支持 ALTER TABLE ADD COLUMN IF NOT EXISTS
# 默认提示词
_GENERATE_PROMPT_DEFAULT = """你是一位幽默大师,专门创作轻松搞笑的短笑话。
{context}
要求:
1. 根据场景和关键词创作一条原创笑话
2. 笑话要有反转或意外结局
3. 语言风格:{style}
4. 字数要求:{length_requirement}
5. 直接输出笑话内容,不需要解释
请严格按照以下 JSON 格式输出,不要加任何额外说明:
{{
"title": "笑话标题",
"content": "笑话正文",
"score": 8,
"reason": "这个笑话巧妙结合了场景和关键词,结尾有反转"
}}
score 是 1-10 的整数评分,reason 是用一句话说明亮点。"""
_OPTIMIZER_QUALITY_DEFAULT = """你是一个幽默内容审核专家。判断以下内容是否是一个合格的笑话/段子。
合格标准(满足任一即可):
1. 有明确的笑点或反转(punchline)
2. 有幽默的语言表达或双关
3. 有意外结局或情理之中意料之外
不合格标准(符合任一即判定不合格):
1. 纯粹的事实陈述,没有任何幽默元素
2. 只是对话片段,没有笑点
3. 普通故事或叙事,没有幽默设计
4. 说教或道理阐述
5. 内容不完整或难以理解
始终返回 JSON 格式:{"has_punchline": true/false, "reason": "简要说明判断理由"}"""
_OPTIMIZER_POLISH_DEFAULT = """你是一个专业的幽默文案编辑。请润色以下笑话,要求:
1. 保持核心笑点不变
2. 优化语言表达,使其更通顺、更精炼
3. 增强节奏感和幽默效果,但不改变原意
4. 字数控制在原内容的 80%-120%
5. 不要添加额外解释或评论
6. 直接输出润色后的内容,不要加任何前缀"""
_OPTIMIZER_EVALUATE_DEFAULT = """你是一个笑话分类和评价专家。对给定的笑话进行分析,返回 JSON 格式的分类和评分结果。
要求:
1. types: 从提供的类型列表中选择所有匹配的类型名称(数组,可以选多个)
2. crowds: 从提供的人群列表中选择所有匹配的人群名称(数组,可以选多个)
3. score: 1-10 分,基于幽默程度、创意和表达效果
4. comment: 简短评语(10字以内)
始终返回 JSON 格式。"""
_CRAWLER_EXTRACT_DEFAULT = """你是一个笑话提取专家。从给定的网页文本中识别并提取所有笑话、幽默段子或有趣内容。
要求:
1. 只返回真正的笑话内容,不要提取普通文章或新闻
2. 每条笑话需要包含:title(简短标题)、content(完整笑话内容)、type(类型)、crowd(人群)
3. 如果网页中没有笑话,返回空数组 []
4. 永远返回合法的 JSON 格式,根节点为数组或包含 jokes 键的对象"""
_CRAWLER_REWRITE_DEFAULT = """你是一个幽默作家,负责润色和改写笑话。
要求:
1. 保持笑话的核心笑点不变
2. 语言更通顺、更幽默
3. 字数控制在原内容的 80%-120% 之间
4. 不要添加任何解释说明"""
def _migrate_db(): def _migrate_db():
from sqlalchemy import inspect, text from sqlalchemy import inspect, text
inspector = inspect(engine) inspector = inspect(engine)
columns = [c["name"] for c in inspector.get_columns("jokes")] columns = [c["name"] for c in inspector.get_columns("ai_settings")]
if "dislike_count" not in columns:
with engine.connect() as conn: # 新增提示词字段
conn.execute(text("ALTER TABLE jokes ADD COLUMN dislike_count INTEGER DEFAULT 0")) prompt_fields = {
conn.commit() "generate_prompt": _GENERATE_PROMPT_DEFAULT,
"optimizer_quality_prompt": _OPTIMIZER_QUALITY_DEFAULT,
"optimizer_polish_prompt": _OPTIMIZER_POLISH_DEFAULT,
"optimizer_evaluate_prompt": _OPTIMIZER_EVALUATE_DEFAULT,
"crawler_extract_prompt": _CRAWLER_EXTRACT_DEFAULT,
"crawler_rewrite_prompt": _CRAWLER_REWRITE_DEFAULT,
}
for field, default_val in prompt_fields.items():
if field not in columns:
from sqlalchemy import Text
col_type = "TEXT"
with engine.connect() as conn:
conn.execute(text(f"ALTER TABLE ai_settings ADD COLUMN {field} {col_type}"))
conn.commit()
# 新增 temperature 字段
temp_fields = [
("optimizer_quality_temperature", "FLOAT", "0.3"),
("optimizer_polish_temperature", "FLOAT", "0.8"),
("optimizer_evaluate_temperature", "FLOAT", "0.3"),
]
for field, col_type, default_val in temp_fields:
if field not in columns:
with engine.connect() as conn:
conn.execute(text(f"ALTER TABLE ai_settings ADD COLUMN {field} {col_type} DEFAULT {default_val}"))
conn.commit()
# 给新增字段填充默认值(已有记录)
with engine.connect() as conn:
for field, default_val in prompt_fields.items():
conn.execute(text(f"UPDATE ai_settings SET {field} = :val WHERE {field} IS NULL"),
{"val": default_val})
conn.commit()
_migrate_db() _migrate_db()
+80 -20
View File
@@ -1,35 +1,90 @@
"""LLM 处理:调用 NVIDIA NIMOpenAI 兼容 API)进行笑话提取、改写和分类""" """LLM 处理:调用 NVIDIA NIMOpenAI 兼容 API)进行笑话提取、改写。"""
import json import json
import os import httpx
from openai import OpenAI from openai import OpenAI
class AiService: class AiService:
def __init__(self, api_base: str, api_key: str, model_name: str, temperature: float = 0.7, max_tokens: int = 2048): def __init__(self, api_base: str, username: str, password: str):
self.client = OpenAI(base_url=api_base, api_key=api_key) self.api_base = api_base.rstrip("/")
self.model = model_name self.username = username
self.temperature = temperature self.password = password
self.max_tokens = max_tokens self.token = None
self.client = None
self.model = ""
self.ai_config = None
def _login(self) -> str:
resp = httpx.post(
f"{self.api_base}/api/auth/login",
json={"username": self.username, "password": self.password},
timeout=30,
)
resp.raise_for_status()
return resp.json()["access_token"]
def _get(self, path: str) -> dict | list:
resp = httpx.get(
f"{self.api_base}{path}",
headers={"Authorization": f"Bearer {self.token}"},
timeout=30,
)
resp.raise_for_status()
return resp.json()
def setup(self):
"""登录并从 API 获取 AI 配置(含提示词和温度)"""
self.token = self._login()
self.ai_config = self._get("/api/admin/settings/active")
self.client = OpenAI(
base_url=self.ai_config["api_base"],
api_key=self.ai_config["api_key"],
)
self.model = self.ai_config["model_name"]
def _get_prompt(self, key: str, default: str) -> str:
if self.ai_config:
val = self.ai_config.get(key)
if val and val.strip():
return val
return default
def extract_jokes(self, page_content: str, known_types: list[str], known_crowds: list[str]) -> list[dict]: def extract_jokes(self, page_content: str, known_types: list[str], known_crowds: list[str]) -> list[dict]:
"""从页面内容中提取笑话,返回结构化数据。""" """从页面内容中提取笑话,返回结构化数据。"""
from crawler.prompts import EXTRACTION_SYSTEM_PROMPT, EXTRACTION_USER_PROMPT system_prompt = self._get_prompt("crawler_extract_prompt",
"""你是一个笑话提取专家。从给定的网页文本中识别并提取所有笑话、幽默段子或有趣内容。
要求:
1. 只返回真正的笑话内容,不要提取普通文章或新闻
2. 每条笑话需要包含:title(简短标题)、content(完整笑话内容)、type(类型)、crowd(人群)
3. 如果网页中没有笑话,返回空数组 []
4. 永远返回合法的 JSON 格式,根节点为数组或包含 jokes 键的对象""")
user_prompt = EXTRACTION_USER_PROMPT.format( user_prompt = f"""网页内容:
page_content=page_content[:8000], ---
known_types=", ".join(known_types), {page_content[:8000]}
known_crowds=", ".join(known_crowds), ---
)
已知笑话类型:{', '.join(known_types)}
已知人群分类:{', '.join(known_crowds)}
请提取所有笑话,以 JSON 格式返回,示例:
[
{{"title": "程序员的幽默", "content": "程序员去相亲...", "types": ["谐音梗", "段子"], "crowds": ["职场", "大学生"]}},
{{"title": "...", "content": "...", "types": ["..."], "crowds": ["..."]}}
]
注意:types 和 crowds 是数组,可以填多个。
只返回 JSON,不要其他文字。"""
response = self.client.chat.completions.create( response = self.client.chat.completions.create(
model=self.model, model=self.model,
messages=[ messages=[
{"role": "system", "content": EXTRACTION_SYSTEM_PROMPT}, {"role": "system", "content": system_prompt},
{"role": "user", "content": user_prompt}, {"role": "user", "content": user_prompt},
], ],
temperature=self.temperature, temperature=self.ai_config.get("temperature", 0.7) if self.ai_config else 0.7,
max_tokens=self.max_tokens, max_tokens=self.ai_config.get("max_tokens", 2048) if self.ai_config else 2048,
) )
raw = response.choices[0].message.content raw = response.choices[0].message.content
@@ -37,13 +92,19 @@ class AiService:
def rewrite_joke(self, content: str) -> str: def rewrite_joke(self, content: str) -> str:
"""润色单条笑话内容。""" """润色单条笑话内容。"""
from crawler.prompts import REWRITE_SYSTEM_PROMPT, REWRITE_USER_PROMPT system_prompt = self._get_prompt("crawler_rewrite_prompt",
"""你是一个幽默作家,负责润色和改写笑话。
要求:
1. 保持笑话的核心笑点不变
2. 语言更通顺、更幽默
3. 字数控制在原内容的 80%-120% 之间
4. 不要添加任何解释说明""")
response = self.client.chat.completions.create( response = self.client.chat.completions.create(
model=self.model, model=self.model,
messages=[ messages=[
{"role": "system", "content": REWRITE_SYSTEM_PROMPT}, {"role": "system", "content": system_prompt},
{"role": "user", "content": REWRITE_USER_PROMPT.format(content=content)}, {"role": "user", "content": f"请润色以下笑话:\n\n{content}"},
], ],
temperature=0.8, temperature=0.8,
max_tokens=500, max_tokens=500,
@@ -60,7 +121,6 @@ class AiService:
jokes = data jokes = data
return [j for j in jokes if isinstance(j, dict) and j.get("content")] return [j for j in jokes if isinstance(j, dict) and j.get("content")]
except json.JSONDecodeError: except json.JSONDecodeError:
# 尝试提取 markdown 代码块
if "```json" in raw: if "```json" in raw:
raw = raw.split("```json")[1].split("```")[0] raw = raw.split("```json")[1].split("```")[0]
elif "```" in raw: elif "```" in raw:
+5 -7
View File
@@ -57,15 +57,13 @@ class Processor:
self.token = self._login() self.token = self._login()
print("[*] 登录成功") print("[*] 登录成功")
ai_config = self._get("/api/admin/settings/active")
self.ai = AiService( self.ai = AiService(
api_base=ai_config["api_base"], api_base=self.api_base,
api_key=ai_config["api_key"], username=self.username,
model_name=ai_config["model_name"], password=self.password,
temperature=ai_config.get("temperature", 0.7),
max_tokens=ai_config.get("max_tokens", 2048),
) )
print(f"[*] AI 配置: {ai_config['model_name']}") self.ai.setup()
print(f"[*] AI 配置: {self.ai.model}")
self.types = self._get("/api/categories/types") self.types = self._get("/api/categories/types")
self.crowds = self._get("/api/categories/crowds") self.crowds = self._get("/api/categories/crowds")
-41
View File
@@ -1,41 +0,0 @@
"""AI 提示词模板。"""
# ===== 笑话提取 =====
EXTRACTION_SYSTEM_PROMPT = """你是一个笑话提取专家。从给定的网页文本中识别并提取所有笑话、幽默段子或有趣内容。
要求:
1. 只返回真正的笑话内容,不要提取普通文章或新闻
2. 每条笑话需要包含:title(简短标题)、content(完整笑话内容)、type(类型)、crowd(人群)
3. 如果网页中没有笑话,返回空数组 []
4. 永远返回合法的 JSON 格式,根节点为数组或包含 jokes 键的对象"""
EXTRACTION_USER_PROMPT = """网页内容:
---
{page_content}
---
已知笑话类型:{known_types}
已知人群分类:{known_crowds}
请提取所有笑话,以 JSON 格式返回,示例:
[
{{"title": "程序员的幽默", "content": "程序员去相亲...", "types": ["谐音梗", "段子"], "crowds": ["职场", "大学生"]}},
{{"title": "...", "content": "...", "types": ["..."], "crowds": ["..."]}}
]
注意:types 和 crowds 是数组,可以填多个。
只返回 JSON,不要其他文字。"""
# ===== 笑话改写 =====
REWRITE_SYSTEM_PROMPT = """你是一个幽默作家,负责润色和改写笑话。
要求:
1. 保持笑话的核心笑点不变
2. 语言更通顺、更幽默
3. 字数控制在原内容的 80%-120% 之间
4. 不要添加任何解释说明"""
REWRITE_USER_PROMPT = """请润色以下笑话:
{content}
只返回润色后的笑话文字,不要其他内容。"""
+76 -23
View File
@@ -8,15 +8,6 @@ import time
import httpx import httpx
from openai import OpenAI from openai import OpenAI
from optimizer.prompts import (
QUALITY_CHECK_SYSTEM_PROMPT,
QUALITY_CHECK_USER_PROMPT,
POLISH_SYSTEM_PROMPT,
POLISH_USER_PROMPT,
EVALUATE_SYSTEM_PROMPT,
EVALUATE_USER_PROMPT,
)
class Optimizer: class Optimizer:
def __init__(self, api_base: str, username: str, password: str): def __init__(self, api_base: str, username: str, password: str):
@@ -26,6 +17,7 @@ class Optimizer:
self.token = None self.token = None
self.ai_client = None self.ai_client = None
self.model_name = "" self.model_name = ""
self.ai_config = None # 完整 AI 配置(含提示词和温度)
self.types = [] self.types = []
self.crowds = [] self.crowds = []
# 统计 # 统计
@@ -67,6 +59,7 @@ class Optimizer:
print("[*] 登录成功") print("[*] 登录成功")
ai_config = self._get("/api/admin/settings/active") ai_config = self._get("/api/admin/settings/active")
self.ai_config = ai_config
self.ai_client = OpenAI( self.ai_client = OpenAI(
base_url=ai_config["api_base"], base_url=ai_config["api_base"],
api_key=ai_config["api_key"], api_key=ai_config["api_key"],
@@ -78,6 +71,22 @@ class Optimizer:
self.crowds = self._get("/api/categories/crowds") self.crowds = self._get("/api/categories/crowds")
print(f"[*] 分类: {len(self.types)} 种类型, {len(self.crowds)} 种人群") print(f"[*] 分类: {len(self.types)} 种类型, {len(self.crowds)} 种人群")
def _get_prompt(self, key: str, default: str) -> str:
"""从数据库配置中读取提示词,没有则返回默认"""
if self.ai_config:
val = self.ai_config.get(key)
if val and val.strip():
return val
return default
def _get_temp(self, key: str, default: float) -> float:
"""从数据库配置中读取温度"""
if self.ai_config:
val = self.ai_config.get(key)
if val is not None:
return float(val)
return default
# === 读取笑话 === # === 读取笑话 ===
def get_jokes(self, status: str | None = None, limit: int | None = None, def get_jokes(self, status: str | None = None, limit: int | None = None,
ids: list[int] | None = None) -> list[dict]: ids: list[int] | None = None) -> list[dict]:
@@ -120,14 +129,32 @@ class Optimizer:
# === Stage 1: 质量检测 === # === Stage 1: 质量检测 ===
def quality_check(self, content: str) -> dict: def quality_check(self, content: str) -> dict:
"""判断笑话是否有笑点,返回 {"has_punchline": bool, "reason": str}""" """判断笑话是否有笑点,返回 {"has_punchline": bool, "reason": str}"""
system_prompt = self._get_prompt("optimizer_quality_prompt",
"""你是一个幽默内容审核专家。判断以下内容是否是一个合格的笑话/段子。
合格标准(满足任一即可):
1. 有明确的笑点或反转(punchline)
2. 有幽默的语言表达或双关
3. 有意外结局或情理之中意料之外
不合格标准(符合任一即判定不合格):
1. 纯粹的事实陈述,没有任何幽默元素
2. 只是对话片段,没有笑点
3. 普通故事或叙事,没有幽默设计
4. 说教或道理阐述
5. 内容不完整或难以理解
始终返回 JSON 格式:{"has_punchline": true/false, "reason": "简要说明判断理由"}""")
temperature = self._get_temp("optimizer_quality_temperature", 0.3)
def _call(): def _call():
return self.ai_client.chat.completions.create( return self.ai_client.chat.completions.create(
model=self.model_name, model=self.model_name,
messages=[ messages=[
{"role": "system", "content": QUALITY_CHECK_SYSTEM_PROMPT}, {"role": "system", "content": system_prompt},
{"role": "user", "content": QUALITY_CHECK_USER_PROMPT.format(content=content[:2000])}, {"role": "user", "content": f"请判断以下内容是否为合格笑话:\n\n{content[:2000]}\n\n返回 JSON 格式。"},
], ],
temperature=0.3, temperature=temperature,
max_tokens=200, max_tokens=200,
) )
resp = self._safe_api_call(_call) resp = self._safe_api_call(_call)
@@ -137,14 +164,24 @@ class Optimizer:
# === Stage 2: AI 润色 === # === Stage 2: AI 润色 ===
def polish(self, content: str) -> str: def polish(self, content: str) -> str:
"""润色笑话内容""" """润色笑话内容"""
system_prompt = self._get_prompt("optimizer_polish_prompt",
"""你是一个专业的幽默文案编辑。请润色以下笑话,要求:
1. 保持核心笑点不变
2. 优化语言表达,使其更通顺、更精炼
3. 增强节奏感和幽默效果,但不改变原意
4. 字数控制在原内容的 80%-120%
5. 不要添加额外解释或评论
6. 直接输出润色后的内容,不要加任何前缀""")
temperature = self._get_temp("optimizer_polish_temperature", 0.8)
def _call(): def _call():
return self.ai_client.chat.completions.create( return self.ai_client.chat.completions.create(
model=self.model_name, model=self.model_name,
messages=[ messages=[
{"role": "system", "content": POLISH_SYSTEM_PROMPT}, {"role": "system", "content": system_prompt},
{"role": "user", "content": POLISH_USER_PROMPT.format(content=content)}, {"role": "user", "content": f"请润色以下笑话:\n\n{content}"},
], ],
temperature=0.8, temperature=temperature,
max_tokens=1024, max_tokens=1024,
) )
resp = self._safe_api_call(_call) resp = self._safe_api_call(_call)
@@ -156,24 +193,40 @@ class Optimizer:
type_names = [t.get("name", "") for t in self.types] type_names = [t.get("name", "") for t in self.types]
crowd_names = [c.get("name", "") for c in self.crowds] crowd_names = [c.get("name", "") for c in self.crowds]
system_prompt = self._get_prompt("optimizer_evaluate_prompt",
"""你是一个笑话分类和评价专家。对给定的笑话进行分析,返回 JSON 格式的分类和评分结果。
要求:
1. types: 从提供的类型列表中选择所有匹配的类型名称(数组,可以选多个)
2. crowds: 从提供的人群列表中选择所有匹配的人群名称(数组,可以选多个)
3. score: 1-10 分,基于幽默程度、创意和表达效果
4. comment: 简短评语(10字以内)
始终返回 JSON 格式。""")
temperature = self._get_temp("optimizer_evaluate_temperature", 0.3)
user_content = f"""笑话内容:
{content[:2000]}
可选类型:{', '.join(type_names)}
可选人群:{', '.join(crowd_names)}
返回 JSON 格式:{{"types": ["类型1", "类型2"], "crowds": ["人群1", "人群2"], "score": 8, "comment": "简短评语"}}"""
def _call(): def _call():
return self.ai_client.chat.completions.create( return self.ai_client.chat.completions.create(
model=self.model_name, model=self.model_name,
messages=[ messages=[
{"role": "system", "content": EVALUATE_SYSTEM_PROMPT}, {"role": "system", "content": system_prompt},
{"role": "user", "content": EVALUATE_USER_PROMPT.format( {"role": "user", "content": user_content},
content=content[:2000],
known_types=", ".join(type_names),
known_crowds=", ".join(crowd_names),
)},
], ],
temperature=0.3, temperature=temperature,
max_tokens=300, max_tokens=300,
) )
resp = self._safe_api_call(_call) resp = self._safe_api_call(_call)
raw = resp.choices[0].message.content.strip() raw = resp.choices[0].message.content.strip()
result = self._parse_json(raw, {"types": [], "crowds": [], "score": 5, "comment": ""}) result = self._parse_json(raw, {"types": [], "crowds": [], "score": 5, "comment": ""})
# Backward compatibility: if LLM returns old single format, convert to array # Backward compatibility
if isinstance(result.get("types"), str): if isinstance(result.get("types"), str):
result["types"] = [result["types"]] if result["types"] else [] result["types"] = [result["types"]] if result["types"] else []
if isinstance(result.get("crowds"), str): if isinstance(result.get("crowds"), str):
-58
View File
@@ -1,58 +0,0 @@
"""AI 提示词模板 — 笑话质量检测、润色、评价分类。"""
# ===== Stage 1: 质量检测 =====
QUALITY_CHECK_SYSTEM_PROMPT = """你是一个幽默内容审核专家。判断以下内容是否是一个合格的笑话/段子。
合格标准(满足任一即可):
1. 有明确的笑点或反转(punchline)
2. 有幽默的语言表达或双关
3. 有意外结局或情理之中意料之外
不合格标准(符合任一即判定不合格):
1. 纯粹的事实陈述,没有任何幽默元素
2. 只是对话片段,没有笑点
3. 普通故事或叙事,没有幽默设计
4. 说教或道理阐述
5. 内容不完整或难以理解
始终返回 JSON 格式:{"has_punchline": true/false, "reason": "简要说明判断理由"}"""
QUALITY_CHECK_USER_PROMPT = """请判断以下内容是否为合格笑话:
{content}
返回 JSON 格式。"""
# ===== Stage 2: AI 润色 =====
POLISH_SYSTEM_PROMPT = """你是一个专业的幽默文案编辑。请润色以下笑话,要求:
1. 保持核心笑点不变
2. 优化语言表达,使其更通顺、更精炼
3. 增强节奏感和幽默效果,但不改变原意
4. 字数控制在原内容的 80%-120%
5. 不要添加额外解释或评论
6. 直接输出润色后的内容,不要加任何前缀"""
POLISH_USER_PROMPT = """请润色以下笑话:
{content}
只输出润色后的笑话内容。"""
# ===== Stage 3: 评价分类 =====
EVALUATE_SYSTEM_PROMPT = """你是一个笑话分类和评价专家。对给定的笑话进行分析,返回 JSON 格式的分类和评分结果。
要求:
1. types: 从提供的类型列表中选择所有匹配的类型名称(数组,可以选多个)
2. crowds: 从提供的人群列表中选择所有匹配的人群名称(数组,可以选多个)
3. score: 1-10 分,基于幽默程度、创意和表达效果
4. comment: 简短评语(10字以内)
始终返回 JSON 格式。"""
EVALUATE_USER_PROMPT = """笑话内容:
{content}
可选类型:{known_types}
可选人群:{known_crowds}
返回 JSON 格式:{{"types": ["类型1", "类型2"], "crowds": ["人群1", "人群2"], "score": 8, "comment": "简短评语"}}"""