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
+13
View File
@@ -14,6 +14,19 @@ class AiSetting(Base):
temperature = Column(Float, nullable=False, default=0.7)
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_keywords = Column(Text, nullable=True, default="")
max_pages_per_run = Column(Integer, nullable=False, default=3)
+33 -33
View File
@@ -1,8 +1,8 @@
"""智能笑话生成器 API"""
"""AI 笑话生成器 API — 提示词从数据库读取"""
import json
from datetime import datetime
from fastapi import APIRouter, Depends, HTTPException, Query
from fastapi import APIRouter, Depends, HTTPException
from openai import OpenAI
from sqlalchemy.orm import Session
@@ -13,9 +13,31 @@ from app.schemas.joke import GenerateRequest, GenerateResponse
router = APIRouter(prefix="/generate", tags=["生成器"])
STYLE_MAP = {
"cold": "冷幽默 / 无厘头",
"warm": "温馨幽默 / 暖心搞笑",
"twist": "反转 / 神转折",
"pun": "谐音梗 / 文字游戏",
"sketch": "段子 / 吐槽调侃",
"irony": "讽刺幽默 / 黑色幽默",
}
# AI Prompt
GENERATION_PROMPT = """你是一位幽默大师,专门创作轻松搞笑的短笑话。
LENGTH_MAP = {
"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}
@@ -36,27 +58,13 @@ GENERATION_PROMPT = """你是一位幽默大师,专门创作轻松搞笑的短
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(
scenarios: list[str],
keywords: list[str],
style: str = "twist",
length: str = "medium",
setting: AiSetting | None = None,
) -> str:
"""构建 AI prompt,支持风格和长度控制"""
parts = []
@@ -66,7 +74,9 @@ def _build_prompt(
parts.append(f"关键词:{', '.join(keywords)}")
if not parts:
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),
style=STYLE_MAP.get(style, "反转 / 神转折"),
length_requirement=LENGTH_MAP.get(length, "80-150字,正常长度"),
@@ -78,7 +88,7 @@ def generate_joke(
req: GenerateRequest,
db: Session = Depends(get_db),
):
"""调用 AI 生成笑话,支持风格和长度控制"""
"""调用 AI 生成笑话,提示词从数据库读取"""
# 获取激活的 AI 配置
setting = db.query(AiSetting).filter(AiSetting.is_active == True).first()
if not setting:
@@ -87,10 +97,9 @@ def generate_joke(
if not setting or not setting.api_key:
raise HTTPException(status_code=503, detail="AI 服务未配置,请联系管理员")
# 调用 AI
try:
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(
model=setting.model_name,
@@ -113,22 +122,18 @@ def _parse_json_response(raw: str | None) -> GenerateResponse:
if not raw or not raw.strip():
return GenerateResponse(title="生成的笑话", content="(内容生成失败,请重新生成)", score=0)
# 尝试从 markdown 代码块中提取 JSON
raw = raw.strip()
if raw.startswith("```"):
lines = raw.split("\n")
# 去掉第一行 ```json 和最后一行 ```
if len(lines) >= 3:
raw = "\n".join(lines[1:-1]).strip()
# 移除可能的尾部分号
if raw.endswith(","):
raw = raw[:-1]
try:
data = json.loads(raw)
except json.JSONDecodeError:
# 如果 JSON 解析失败,尝试查找花括号内的内容
try:
start = raw.index("{")
end = raw.rindex("}") + 1
@@ -147,20 +152,17 @@ def _parse_json_response(raw: str | None) -> GenerateResponse:
def _parse_response(raw: str | None) -> GenerateResponse:
"""解析 AI 返回内容,提取标题和内容"""
# 防御:处理空或 None 输入
"""解析 AI 返回内容(非 JSON 回退)"""
if not raw or not raw.strip():
raise ValueError("AI 返回内容为空")
title = ""
content = raw
# 尝试提取 "标题:xxx" 或 "标题:xxx"
for line in raw.split("\n"):
line = line.strip()
if line.startswith("标题:") or line.startswith("标题:"):
title = line.split("", 1)[-1].split(":", 1)[-1].strip()
# 只替换这一行,不要 replace 全局
lines = content.split("\n")
for i, l in enumerate(lines):
if l.strip() == line:
@@ -169,14 +171,12 @@ def _parse_response(raw: str | None) -> GenerateResponse:
content = "\n".join(lines).strip()
break
# 如果没有提取到标题,取第一行
if not title:
first_line = raw.split("\n")[0].strip()
if first_line.startswith("标题"):
first_line = first_line.split("", 1)[-1].split(":", 1)[-1].strip()
title = first_line[:30] if len(first_line) > 30 else first_line
# 防御:content 不能为空
if not content.strip():
content = "(内容生成失败,请重新生成)"
+22 -4
View File
@@ -22,12 +22,12 @@ def list_settings(
@router.get("/active")
def get_active_setting(
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()
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
@@ -93,4 +93,22 @@ def delete_setting(
raise HTTPException(status_code=404, detail="配置不存在")
db.delete(db_setting)
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):
@@ -8,6 +8,20 @@ class AiSettingBase(BaseModel):
model_name: str = Field(default="nvidia/llama-3.1-nemotron-70b-instruct")
temperature: float = Field(default=0.7, ge=0, le=2)
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_keywords: str = Field(default="")
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)
# 对已有表新增字段的兼容迁移(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():
from sqlalchemy import inspect, text
inspector = inspect(engine)
columns = [c["name"] for c in inspector.get_columns("jokes")]
if "dislike_count" not in columns:
with engine.connect() as conn:
conn.execute(text("ALTER TABLE jokes ADD COLUMN dislike_count INTEGER DEFAULT 0"))
conn.commit()
columns = [c["name"] for c in inspector.get_columns("ai_settings")]
# 新增提示词字段
prompt_fields = {
"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()