485 lines
20 KiB
Python
485 lines
20 KiB
Python
#!/usr/bin/env python3
|
||
"""
|
||
平台适配文章撰写
|
||
根据大纲和平台配置(字数/格式/配图要求),为知乎/公众号/小红书各平台生成适配内容
|
||
"""
|
||
import json, datetime, logging, sys, re
|
||
from pathlib import Path
|
||
from typing import Dict, List
|
||
|
||
PROJECT_ROOT = Path(__file__).parent.parent
|
||
sys.path.insert(0, str(PROJECT_ROOT))
|
||
sys.path.insert(0, str(PROJECT_ROOT / 'platform' / 'backend'))
|
||
|
||
from db_helper import get_topic_by_id, update_topic_status, save_article
|
||
from content_cleaner import strip_thinking, strip_ai_preface, clean_markdown_content, clean_html_content
|
||
from prompt_loader import get_prompt, get_prompt_params
|
||
from image_generator import insert_lead_image
|
||
try:
|
||
from app.core.nvidia_client import call_llm
|
||
HAVE_LLM = True
|
||
except ImportError:
|
||
HAVE_LLM = False
|
||
|
||
import mistune
|
||
|
||
DATA_DIR = PROJECT_ROOT / "automation" / "data"
|
||
OUTLINE_DIR = DATA_DIR / "outlines"
|
||
TEMPLATES_DIR = PROJECT_ROOT / "automation" / "templates"
|
||
LOGS_DIR = PROJECT_ROOT / "automation" / "logs"
|
||
TODAY = datetime.datetime.now().strftime("%Y-%m-%d")
|
||
GEN_TIME = datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S")
|
||
|
||
logging.basicConfig(
|
||
level=logging.INFO,
|
||
format='%(asctime)s - %(levelname)s - %(message)s',
|
||
handlers=[
|
||
logging.FileHandler(LOGS_DIR / f"writer_{TODAY}.log"),
|
||
logging.StreamHandler()
|
||
]
|
||
)
|
||
logger = logging.getLogger(__name__)
|
||
|
||
_md_parser = mistune.create_markdown()
|
||
|
||
_FALLBACK_PLATFORM_CONFIG = {
|
||
"zhihu": {"max_chars": 3000, "style": "深度长文分析", "min_chars": 1500},
|
||
"wechat": {"max_chars": 1500, "style": "亲切口语化", "min_chars": 800},
|
||
"xiaohongshu": {"max_chars": 800, "style": "图文笔记,emoji+标签", "min_chars": 300},
|
||
}
|
||
|
||
def _load_platform_config() -> dict:
|
||
try:
|
||
from app.database import SessionLocal
|
||
from app.models import PlatformConfig
|
||
db = SessionLocal()
|
||
configs = db.query(PlatformConfig).all()
|
||
db.close()
|
||
result = {}
|
||
for c in configs:
|
||
result[c.platform] = {
|
||
"max_chars": c.max_words or _FALLBACK_PLATFORM_CONFIG.get(c.platform, {}).get("max_chars", 3000),
|
||
"style": c.default_format or _FALLBACK_PLATFORM_CONFIG.get(c.platform, {}).get("style", "深度内容"),
|
||
"min_chars": c.min_words or _FALLBACK_PLATFORM_CONFIG.get(c.platform, {}).get("min_chars", 300),
|
||
}
|
||
return result
|
||
except Exception:
|
||
pass
|
||
return dict(_FALLBACK_PLATFORM_CONFIG)
|
||
|
||
PLATFORM_CONFIG = _load_platform_config()
|
||
|
||
class Writer:
|
||
def __init__(self, topic_id: str):
|
||
self.topic_id = topic_id
|
||
self.topic = self._load_topic()
|
||
outline_file = OUTLINE_DIR / TODAY / f"{topic_id}_outline.md"
|
||
if not outline_file.exists():
|
||
raise FileNotFoundError(f"Outline not found: {outline_file}")
|
||
self.outline_content = outline_file.read_text(encoding='utf-8')
|
||
research_file = DATA_DIR / "research" / TODAY / f"{topic_id}_research.md"
|
||
self.research_notes = research_file.read_text(encoding='utf-8') if research_file.exists() else ""
|
||
|
||
def _load_topic(self) -> Dict:
|
||
topic = get_topic_by_id(self.topic_id)
|
||
if not topic:
|
||
raise ValueError(f"Topic {self.topic_id} not found")
|
||
return topic
|
||
|
||
def _clean_title(self, title: str) -> str:
|
||
title = re.sub(r'[((]约\s*\d+字[))]', '', title)
|
||
title = re.sub(r'[((]MVP[))]', '', title)
|
||
title = re.sub(r'[((][^))]*?[))]', '', title)
|
||
return title.strip()
|
||
|
||
def _parse_outline_sections(self) -> List[Dict]:
|
||
sections = []
|
||
current = None
|
||
for line in self.outline_content.splitlines():
|
||
if line.startswith("# "):
|
||
if current:
|
||
sections.append(current)
|
||
current = {"level": 1, "title": line[2:].strip(), "content": "", "section_type": "normal"}
|
||
elif line.startswith("## "):
|
||
if current:
|
||
sections.append(current)
|
||
title = line[3:].strip()
|
||
stype = "noise" if title in ("文章大纲", "大纲", "文章结构", "结构") else "normal"
|
||
current = {"level": 2, "title": title, "content": "", "section_type": stype}
|
||
elif line.startswith("### "):
|
||
if current:
|
||
sections.append(current)
|
||
current = {"level": 3, "title": line[4:].strip(), "content": "", "section_type": "normal"}
|
||
else:
|
||
if current and line.strip():
|
||
current['content'] = current.get('content', '') + line + "\n"
|
||
if current:
|
||
sections.append(current)
|
||
return sections
|
||
|
||
@staticmethod
|
||
def _clean_markdown(text: str) -> str:
|
||
return clean_markdown_content(text)
|
||
|
||
@staticmethod
|
||
def _is_outline_noise(line: str) -> bool:
|
||
stripped = line.strip()
|
||
if not stripped:
|
||
return True
|
||
if stripped.startswith('---'):
|
||
return True
|
||
if '大纲生成时间' in stripped:
|
||
return True
|
||
if stripped.startswith('*大纲'):
|
||
return True
|
||
if re.match(r'^\*{0,3}【[^】]*】\*{0,3}\s*$', stripped):
|
||
return True
|
||
return False
|
||
|
||
def _is_bullet_only(self, text: str) -> bool:
|
||
"""检查内容是否主要是要点列表(大纲格式),需要 LLM 展开"""
|
||
lines = [l.strip() for l in text.split('\n') if l.strip()]
|
||
if not lines:
|
||
return False
|
||
bullet_count = sum(1 for l in lines if l.startswith(('- ', '* ', '**', '+ ')))
|
||
return bullet_count / len(lines) > 0.4
|
||
|
||
def _expand_section(self, section: Dict) -> str:
|
||
content = section.get('content', '').strip()
|
||
# 大纲要点格式(>40% 行以 -/*/** 开头)或内容过短(<100字)应由 LLM 展开为连贯段落
|
||
if HAVE_LLM and (self._is_bullet_only(content) or (content and len(content) < 100)):
|
||
logger.info(f"使用 LLM 扩写章节(要点→段落): {section['title']}")
|
||
prompt = get_prompt("section_expansion",
|
||
topic_title=self.topic['title'],
|
||
section_title=section['title'],
|
||
date=datetime.datetime.now().strftime('%Y年%m月%d日'),
|
||
content=content,
|
||
)
|
||
try:
|
||
params = get_prompt_params("section_expansion")
|
||
expanded = call_llm(prompt, temperature=params.get("temperature", 0.75), max_tokens=params.get("max_tokens", 3000))
|
||
if expanded:
|
||
cleaned = self._clean_markdown(expanded.strip())
|
||
if cleaned:
|
||
return cleaned
|
||
except Exception as e:
|
||
logger.warning(f"LLM 扩写失败: {e}")
|
||
|
||
# Fallback: 将 bullet points 展开为段落(过滤噪音行)
|
||
lines = [l.strip() for l in content.split('\n') if not self._is_outline_noise(l)]
|
||
if lines:
|
||
sentences = []
|
||
for line in lines:
|
||
text = line
|
||
for prefix in ['- ', '* ', '1. ', '2. ', '3. ', '4. ', '5. ']:
|
||
if line.startswith(prefix):
|
||
text = line[len(prefix):]
|
||
break
|
||
text = text.strip()
|
||
if text:
|
||
if text[-1] not in '。!?;':
|
||
text += '。'
|
||
sentences.append(text)
|
||
if sentences:
|
||
result = ' '.join(sentences)
|
||
return self._clean_markdown(result)
|
||
return ''
|
||
|
||
def generate_full_markdown(self) -> str:
|
||
sections = self._parse_outline_sections()
|
||
parts = []
|
||
for sec in sections:
|
||
if sec['level'] == 1:
|
||
if sec.get('content'):
|
||
expanded = self._expand_section(sec)
|
||
if expanded:
|
||
parts.append(expanded + "\n")
|
||
continue
|
||
# 跳过大纲结构噪音节点
|
||
title_stripped = sec['title'].strip()
|
||
if title_stripped in ('文章大纲', '大纲', '文章结构', '结构'):
|
||
continue
|
||
if sec.get('section_type') == 'noise':
|
||
continue
|
||
heading = f"{'#' * sec['level']} {sec['title']}"
|
||
parts.append(heading)
|
||
if sec.get('content'):
|
||
expanded = self._expand_section(sec)
|
||
parts.append(expanded + "\n")
|
||
full_md = "\n".join(parts).strip()
|
||
# 收集所有引用来源,统一添加到文末
|
||
refs = set()
|
||
for m in re.finditer(r'(来源:([^)]+))', full_md):
|
||
refs.add(m.group(1).strip())
|
||
if refs:
|
||
# 去掉已有参考资料区,重新生成统一的
|
||
full_md = re.sub(r'\n---\n\*\*参考资料\*\*[\s\S]*$', '', full_md).strip()
|
||
ref_lines = "\n".join(f"- {r}" for r in sorted(refs))
|
||
full_md += f"\n\n---\n\n**参考资料**\n{ref_lines}"
|
||
return full_md
|
||
|
||
def _adapt_for_platform(self, markdown: str, platform: str) -> str:
|
||
cfg = PLATFORM_CONFIG[platform]
|
||
max_c = cfg['max_chars']
|
||
lines = markdown.split('\n')
|
||
|
||
if platform == "zhihu":
|
||
result = []
|
||
in_list = False
|
||
for line in lines:
|
||
stripped = line.strip()
|
||
# 数据类行 → 引用格式(知乎文章中引用数据能增强可信度)
|
||
if any(stripped.startswith(p) for p in ('据统计', '调研显示', '数据显示', '报告指出', '根据', '数据显示')):
|
||
line = f"> {line}"
|
||
# 列表保持原样(知乎支持 markdown 列表)
|
||
if stripped.startswith('- ') or stripped.startswith('* '):
|
||
if not in_list:
|
||
result.append('')
|
||
in_list = True
|
||
else:
|
||
in_list = False
|
||
result.append(line)
|
||
adapted = '\n'.join(result)
|
||
# 末尾加讨论引导(知乎算法权重:互动率)
|
||
if not any(kw in adapted for kw in ('你觉得', '你怎么看', '欢迎在评论区', '说说你的')):
|
||
adapted += "\n\n---\n\n你觉得这个观点有道理吗?你在工作中有没有类似的经验?欢迎在评论区聊聊。"
|
||
return adapted
|
||
|
||
if platform == "xiaohongshu":
|
||
result = []
|
||
char_count = 0
|
||
last_was_heading = False
|
||
for line in lines:
|
||
if char_count >= max_c:
|
||
break
|
||
stripped = line.strip()
|
||
if line.startswith('## '):
|
||
if not last_was_heading and result:
|
||
result.append('')
|
||
char_count += 1
|
||
line = f"## ✨ {line[3:]}"
|
||
last_was_heading = True
|
||
elif line.startswith('### '):
|
||
if not last_was_heading and result:
|
||
result.append('')
|
||
char_count += 1
|
||
line = f"### 💡 {line[4:]}"
|
||
last_was_heading = True
|
||
else:
|
||
last_was_heading = False
|
||
# 超长段落后拆行 + 每段前加点缀
|
||
if stripped and len(stripped) > 60:
|
||
sentences = [s.strip() for s in stripped.replace('。', '。\n').split('\n') if s.strip()]
|
||
for s in sentences:
|
||
if s and char_count < max_c:
|
||
result.append(s)
|
||
char_count += len(s)
|
||
continue
|
||
result.append(line)
|
||
char_count += len(line)
|
||
adapted = '\n'.join(result)
|
||
# 结尾加收藏引导(小红书算法权重:收藏率)
|
||
if '收藏' not in adapted:
|
||
adapted += "\n\n✨ 觉得有用的话点个收藏吧,下次需要的时候随时翻出来看~"
|
||
return adapted
|
||
|
||
if platform == "wechat":
|
||
result = []
|
||
for line in lines:
|
||
# 人称统一:我们→我,你们→你
|
||
line = line.replace('我们', '我').replace('你们', '你').replace('我', '你')
|
||
if line.startswith('### '):
|
||
result.append(f"\n**{line[4:]}**\n")
|
||
elif line.startswith('## '):
|
||
result.append(f"\n**{line[3:]}**\n")
|
||
elif line.strip() and len(line) > 80:
|
||
sentences = [s.strip() for s in line.replace('。', '。\n').split('\n') if s.strip()]
|
||
for s in sentences:
|
||
if s:
|
||
result.append(s)
|
||
else:
|
||
result.append(line)
|
||
adapted = '\n'.join(result)
|
||
return adapted
|
||
|
||
return markdown
|
||
|
||
def _get_platform_tags(self, platform: str) -> str:
|
||
field = self.topic.get('field', '')
|
||
title = self.topic.get('title', '')
|
||
core = self.topic.get('core_concept', '')
|
||
|
||
tag_prompts = {
|
||
"zhihu": get_prompt("tags_generation", platform="知乎", title=title, field=field, core=core),
|
||
"wechat": get_prompt("tags_generation", platform="公众号", title=title, field=field, core=core),
|
||
"xiaohongshu": get_prompt("tags_generation", platform="小红书", title=title, field=field, core=core),
|
||
}
|
||
|
||
if HAVE_LLM:
|
||
prompt = tag_prompts.get(platform, get_prompt("tags_generation", platform=platform, title=title, field=field, core=core))
|
||
try:
|
||
params = get_prompt_params("tags_generation")
|
||
tags_text = call_llm(prompt, temperature=params.get("temperature", 0.3), max_tokens=params.get("max_tokens", 500))
|
||
tags_text = strip_thinking(tags_text)
|
||
if tags_text:
|
||
tags = [t.strip('#') for t in tags_text.strip().split() if t.strip('#')]
|
||
if tags:
|
||
return " ".join(f'<span class="tag">{t}</span>' for t in tags[:5])
|
||
except Exception:
|
||
pass
|
||
|
||
tags = []
|
||
if field:
|
||
import re
|
||
parts = re.split(r'[/、与和及]', field)
|
||
for p in parts:
|
||
p = p.strip()
|
||
if len(p) >= 2:
|
||
tags.append(p)
|
||
if len(parts) == 1 and len(parts[0]) > 4:
|
||
for i in range(0, len(parts[0]), 2):
|
||
chunk = parts[0][i:i+2]
|
||
if len(chunk) == 2:
|
||
tags.append(chunk)
|
||
tags.pop(0)
|
||
|
||
platform_extra = {"zhihu": ["职场"], "xiaohongshu": ["生活"]}
|
||
for t in platform_extra.get(platform, []):
|
||
if t not in tags:
|
||
tags.append(t)
|
||
|
||
if not tags:
|
||
tags = ["科技"]
|
||
|
||
seen = set()
|
||
return " ".join(f'<span class="tag">{t}</span>' for t in tags if t not in seen and not seen.add(t))
|
||
|
||
def _optimize_title(self, platform: str) -> str:
|
||
original = self.topic['title']
|
||
if not HAVE_LLM:
|
||
return original
|
||
|
||
if platform == "zhihu":
|
||
prompt = get_prompt("title_optimize_zhihu",
|
||
title=original,
|
||
core=self.topic.get('core_concept', ''),
|
||
pain=self.topic.get('audience_pain', ''),
|
||
field=self.topic.get('field', ''),
|
||
)
|
||
elif platform == "wechat":
|
||
prompt = get_prompt("title_optimize_wechat",
|
||
title=original,
|
||
core=self.topic.get('core_concept', ''),
|
||
)
|
||
elif platform == "xiaohongshu":
|
||
prompt = get_prompt("title_optimize_xhs",
|
||
title=original,
|
||
core=self.topic.get('core_concept', ''),
|
||
)
|
||
else:
|
||
prompt = f"给以下文章改个吸引人的{platform}标题:{original}"
|
||
|
||
try:
|
||
if platform in ("zhihu", "wechat", "xiaohongshu"):
|
||
params = get_prompt_params(f"title_optimize_{platform}")
|
||
resp = call_llm(prompt, temperature=params.get("temperature", 0.8), max_tokens=params.get("max_tokens", 1500))
|
||
else:
|
||
resp = call_llm(prompt, temperature=0.7)
|
||
resp = strip_thinking(resp)
|
||
titles = []
|
||
for line in resp.strip().split('\n'):
|
||
line = line.strip()
|
||
if not line:
|
||
continue
|
||
line = re.sub(r'^\d+[.、)\s]+', '', line)
|
||
line = line.strip('*#- \t"\'"''"')
|
||
# 跳过思考/建议类输出(如"不如:"、"或者:"、"建议方案"等)
|
||
if re.match(r'^(不如|或者|建议|推荐|参考|方案[一二三]|第[一二三]种|以[下是]|标题[一二三]|选项)', line):
|
||
continue
|
||
if line:
|
||
titles.append(line)
|
||
if titles:
|
||
best = titles[0][:80]
|
||
# 如果优化后标题与原文毫无关联或过短,回退原题
|
||
if len(best) < 4 or (len(set(best) & set(original)) < 2 and len(original) > 4):
|
||
logger.warning(f"标题优化结果异常「{best}」,回退原文")
|
||
return original
|
||
logger.info(f"标题优化 [{platform}]: {best}")
|
||
return best
|
||
except Exception as e:
|
||
logger.warning(f"标题优化失败: {e}")
|
||
return original
|
||
|
||
def generate_platform_html(self, markdown: str, platform: str) -> str:
|
||
title = self._optimize_title(platform)
|
||
adapted = self._adapt_for_platform(markdown, platform)
|
||
tpl_path = TEMPLATES_DIR / f"{platform}.html"
|
||
if tpl_path.exists():
|
||
template = tpl_path.read_text(encoding='utf-8')
|
||
else:
|
||
template = "<!DOCTYPE html><html><head><meta charset='UTF-8'><title>{{TITLE}}</title><meta name='viewport' content='width=device-width'><style>body{max-width:800px;margin:0 auto;padding:20px;font-family:-apple-system,sans-serif;line-height:1.8}</style></head><body><h1>{{TITLE}}</h1><!-- CONTENT --></body></html>"
|
||
|
||
adapted = strip_ai_preface(adapted)
|
||
html = template.replace("{{TITLE}}", title).replace("{{DATE}}", TODAY).replace("{{GEN_TIME}}", GEN_TIME)
|
||
html_content = _md_parser(adapted)
|
||
|
||
# 仅插入头图(每个平台一篇一张,不过度)
|
||
html_content = insert_lead_image(
|
||
html_content, platform,
|
||
title=self.topic.get('title', title),
|
||
field=self.topic.get('field', ''),
|
||
)
|
||
|
||
html = html.replace("<!-- CONTENT -->", html_content)
|
||
|
||
# 防御:清理可能在 LLM 输出中混入的 markdown 代码围栏和文件头
|
||
html = re.sub(r'^```+\w*\s*\n?', '', html)
|
||
html = re.sub(r'\n?```+\s*$', '', html)
|
||
html = html.strip()
|
||
|
||
tags_html = self._get_platform_tags(platform)
|
||
if tags_html:
|
||
html = html.replace("<!-- TAGS -->", tags_html)
|
||
else:
|
||
html = html.replace("<!-- TAGS -->", "")
|
||
|
||
return html
|
||
|
||
def save_html(self, html: str, platform: str) -> str:
|
||
try:
|
||
save_article(self.topic_id, platform, html)
|
||
logger.info(f"文章写入数据库: {platform}_{self.topic_id}")
|
||
return f"db:{platform}_{self.topic_id}"
|
||
except Exception as e:
|
||
logger.warning(f"数据库保存失败: {e}")
|
||
return ""
|
||
|
||
def mark_draft(self):
|
||
update_topic_status(self.topic_id, 'review')
|
||
logger.info(f"选题 {self.topic_id} 状态已更新为待审查(数据库)")
|
||
|
||
def run(self):
|
||
logger.info("开始撰写阶段")
|
||
markdown = self.generate_full_markdown()
|
||
results = {}
|
||
for platform in ["zhihu", "wechat", "xiaohongshu"]:
|
||
html = self.generate_platform_html(markdown, platform)
|
||
results[platform] = str(self.save_html(html, platform))
|
||
self.mark_draft()
|
||
logger.info(f"撰写完成,状态已更新为待审查")
|
||
return {"ok": True, "files": results}
|
||
|
||
def main():
|
||
import argparse
|
||
parser = argparse.ArgumentParser()
|
||
parser.add_argument('--topic-id', required=True, help='选题ID')
|
||
args = parser.parse_args()
|
||
|
||
w = Writer(args.topic_id)
|
||
result = w.run()
|
||
print(json.dumps(result, ensure_ascii=False))
|
||
sys.exit(0 if result['ok'] else 1)
|
||
|
||
if __name__ == "__main__":
|
||
main()
|