890 lines
32 KiB
Python
890 lines
32 KiB
Python
#!/usr/bin/env python3
|
||
"""
|
||
台湾 2028 知识库批量生成工具
|
||
|
||
基于 prompts_master.json + 多个 *_full_list.json,
|
||
调用 OpenAI 格式 API 批量生成知识库文件。
|
||
|
||
依赖:
|
||
pip install openai
|
||
|
||
基本用法:
|
||
# 跑所有清单,输出到 ./knowledge_generated/
|
||
python generate_knowledge.py \\
|
||
--api-key YOUR_KEY \\
|
||
--base-url https://api.openai.com/v1 \\
|
||
--model gpt-4o \\
|
||
--master-prompt prompts_master.json \\
|
||
--lists p1_full_list.json p2_full_list.json p3_full_list.json p_mil_ext_full_list.json \\
|
||
--reference-dir ./knowledge_existing \\
|
||
--output-dir ./knowledge_generated \\
|
||
--summary-file ./run_summary.json
|
||
|
||
# 干跑(不调用 API),看会跑哪些 item
|
||
python generate_knowledge.py ... --dry-run
|
||
|
||
# 只跑特定优先级
|
||
python generate_knowledge.py ... --priority P1
|
||
|
||
# 只跑特定议题组
|
||
python generate_knowledge.py ... --topic-group taiwan_military
|
||
|
||
# 只跑特定 ID(支持多个)—— 用于重跑失败的
|
||
python generate_knowledge.py ... --ids P1-MIL-PERSON-01 P1-MIL-ORG-01
|
||
|
||
# 限制数量(先测试 3 份)
|
||
python generate_knowledge.py ... --limit 3
|
||
|
||
# 并发执行(默认 1 顺序)
|
||
python generate_knowledge.py ... --concurrency 3
|
||
|
||
# 兼容多种 OpenAI 格式 API:
|
||
# OpenAI: --base-url https://api.openai.com/v1 --model gpt-4o
|
||
# DeepSeek: --base-url https://api.deepseek.com --model deepseek-chat
|
||
# 阿里通义: --base-url https://dashscope.aliyuncs.com/compatible-mode/v1 --model qwen-max
|
||
# 本地 vLLM: --base-url http://localhost:8000/v1 --model your-local-model
|
||
# Ollama: --base-url http://localhost:11434/v1 --model llama3.1
|
||
"""
|
||
|
||
import argparse
|
||
import asyncio
|
||
import json
|
||
import logging
|
||
import re
|
||
import sys
|
||
import time
|
||
from dataclasses import dataclass, field
|
||
from pathlib import Path
|
||
from typing import Any, Optional
|
||
|
||
|
||
# ============================================================
|
||
# 数据结构
|
||
# ============================================================
|
||
|
||
@dataclass
|
||
class GenerationItem:
|
||
"""单一待生成文件"""
|
||
id: str
|
||
path: str
|
||
name: str
|
||
type: str
|
||
priority: str
|
||
topic_group: str
|
||
topic: str
|
||
style_reference: list[str] = field(default_factory=list)
|
||
must_cover: list[str] = field(default_factory=list)
|
||
key_relations: dict[str, str] = field(default_factory=dict)
|
||
notes: str = ""
|
||
source_list: str = "" # 来源清单文件名
|
||
|
||
@classmethod
|
||
def from_dict(cls, data: dict, source_list: str = "") -> "GenerationItem":
|
||
return cls(
|
||
id=data.get("id", ""),
|
||
path=data.get("path", ""),
|
||
name=data.get("name", ""),
|
||
type=data.get("type", ""),
|
||
priority=data.get("priority", ""),
|
||
topic_group=data.get("topic_group", ""),
|
||
topic=data.get("topic", ""),
|
||
style_reference=data.get("style_reference", []),
|
||
must_cover=data.get("must_cover", []),
|
||
key_relations=data.get("key_relations", {}),
|
||
notes=data.get("notes", ""),
|
||
source_list=source_list,
|
||
)
|
||
|
||
|
||
@dataclass
|
||
class GenerationResult:
|
||
"""单一文件的生成结果"""
|
||
item: GenerationItem
|
||
success: bool
|
||
content: Optional[str] = None
|
||
error: Optional[str] = None
|
||
elapsed_seconds: float = 0.0
|
||
input_tokens: int = 0
|
||
output_tokens: int = 0
|
||
|
||
|
||
# ============================================================
|
||
# 文件加载
|
||
# ============================================================
|
||
|
||
def load_master_prompt(path: Path) -> dict:
|
||
"""加载通用提示词文件"""
|
||
if not path.exists():
|
||
sys.exit(f"错误: 找不到通用提示词文件 {path}")
|
||
with path.open(encoding="utf-8") as f:
|
||
return json.load(f)
|
||
|
||
|
||
def load_list_files(paths: list[Path]) -> list[GenerationItem]:
|
||
"""加载所有清单文件,合并 items"""
|
||
all_items: list[GenerationItem] = []
|
||
for path in paths:
|
||
if not path.exists():
|
||
logging.warning(f"清单文件不存在,跳过: {path}")
|
||
continue
|
||
with path.open(encoding="utf-8") as f:
|
||
data = json.load(f)
|
||
items_raw = data.get("items", [])
|
||
for item_data in items_raw:
|
||
item = GenerationItem.from_dict(item_data, source_list=path.name)
|
||
all_items.append(item)
|
||
logging.info(f"加载清单 {path.name}:{len(items_raw)} 个 items")
|
||
return all_items
|
||
|
||
|
||
def load_reference_file(reference_dir: Path, rel_path: str) -> Optional[str]:
|
||
"""加载风格参考文件"""
|
||
full_path = reference_dir / rel_path
|
||
if not full_path.exists():
|
||
logging.warning(f" 风格参考文件不存在: {full_path}")
|
||
return None
|
||
try:
|
||
return full_path.read_text(encoding="utf-8")
|
||
except Exception as e:
|
||
logging.warning(f" 无法读取风格参考文件 {full_path}: {e}")
|
||
return None
|
||
|
||
|
||
# ============================================================
|
||
# Prompt 组装
|
||
# ============================================================
|
||
|
||
def build_system_prompt(master: dict) -> str:
|
||
"""组装 system prompt——通用部分"""
|
||
mp = master["master_prompt"]
|
||
|
||
parts = [
|
||
"# 角色",
|
||
mp["role"],
|
||
"",
|
||
"# 写作哲学",
|
||
"",
|
||
"## 核心原则",
|
||
mp["writing_philosophy"]["core_principle"],
|
||
"",
|
||
"## 避免",
|
||
]
|
||
for x in mp["writing_philosophy"]["what_to_avoid"]:
|
||
parts.append(f"- {x}")
|
||
|
||
parts.extend(["", "## 追求"])
|
||
for x in mp["writing_philosophy"]["what_to_pursue"]:
|
||
parts.append(f"- {x}")
|
||
|
||
parts.extend([
|
||
"",
|
||
"# 风格要求",
|
||
f"- 语言:{mp['stylistic_requirements']['language']}",
|
||
f"- 语气:{mp['stylistic_requirements']['tone']}",
|
||
f"- 用词:{mp['stylistic_requirements']['wording']}",
|
||
f"- 格式:{mp['stylistic_requirements']['format']}",
|
||
"",
|
||
"# 必含段落",
|
||
])
|
||
for section_name, section_items in mp["mandatory_sections"].items():
|
||
parts.append(f"\n## {section_name}")
|
||
for x in section_items:
|
||
parts.append(f"- {x}")
|
||
|
||
parts.extend(["", "# 深度指标 —— 好的档案应有"])
|
||
for x in mp["depth_indicators"]["what_a_good_file_has"]:
|
||
parts.append(f"- {x}")
|
||
|
||
parts.extend(["", "# 反模式 —— 严格避免"])
|
||
for x in mp["anti_patterns"]:
|
||
parts.append(f"- {x}")
|
||
|
||
return "\n".join(parts)
|
||
|
||
|
||
def build_type_prompt(master: dict, item_type: str) -> str:
|
||
"""组装类型专属 prompt"""
|
||
type_prompts = master["type_prompts"]
|
||
if item_type not in type_prompts:
|
||
logging.warning(f" 未知类型 '{item_type}',使用通用规范")
|
||
return ""
|
||
|
||
tp = type_prompts[item_type]
|
||
parts = [
|
||
f"# 类型规范:{item_type}",
|
||
"",
|
||
f"## 描述",
|
||
tp["description"],
|
||
"",
|
||
f"## 字数区间",
|
||
tp["length_range"],
|
||
"",
|
||
"## 必须涵盖的维度",
|
||
]
|
||
for x in tp["must_cover_dimensions"]:
|
||
parts.append(f"- {x}")
|
||
|
||
parts.extend(["", "## 写作注意事项"])
|
||
for x in tp.get("additional_notes", []):
|
||
parts.append(f"- {x}")
|
||
|
||
return "\n".join(parts)
|
||
|
||
|
||
def build_user_prompt(
|
||
item: GenerationItem,
|
||
reference_contents: list[tuple[str, str]],
|
||
) -> str:
|
||
"""组装 user prompt——本份文件的具体指令"""
|
||
parts = [
|
||
"# 本次任务:撰写以下档案",
|
||
"",
|
||
f"**文件路径**: `{item.path}`",
|
||
f"**名称**: {item.name}",
|
||
f"**类型**: {item.type}",
|
||
f"**议题分组**: {item.topic_group}",
|
||
f"**优先级**: {item.priority}",
|
||
"",
|
||
f"## 主题",
|
||
item.topic,
|
||
"",
|
||
"## 必须涵盖的具体维度",
|
||
]
|
||
for x in item.must_cover:
|
||
parts.append(f"- {x}")
|
||
|
||
parts.extend(["", "## 关键关联文件(在文中适当位置交叉引用)"])
|
||
for rel_path, function in item.key_relations.items():
|
||
parts.append(f"- `{rel_path}` —— {function}")
|
||
|
||
if item.notes:
|
||
parts.extend(["", "## 本档案的特定写作注意事项", item.notes])
|
||
|
||
if reference_contents:
|
||
parts.extend([
|
||
"",
|
||
"---",
|
||
"",
|
||
"# 风格参考——以下是已写好的同类档案,请参照其结构、深度、用词、分析方式",
|
||
])
|
||
for ref_path, content in reference_contents:
|
||
parts.extend([
|
||
"",
|
||
f"## 参考档案:`{ref_path}`",
|
||
"",
|
||
"```markdown",
|
||
content,
|
||
"```",
|
||
])
|
||
|
||
parts.extend([
|
||
"",
|
||
"---",
|
||
"",
|
||
"# 输出要求",
|
||
"",
|
||
"1. **只输出最终的 Markdown 档案内容**——不要任何前言、说明、代码块包裹。",
|
||
"2. 直接以档案的第一行(标题行 `# XXX`)开始,以最后一行(`last_updated:` 标注)结束。",
|
||
"3. 严格遵守类型规范的字数区间。",
|
||
"4. 必须涵盖『必须涵盖的具体维度』中所有项目。",
|
||
"5. 在文中适当位置交叉引用『关键关联文件』中至少 70% 的文件。",
|
||
"6. 遵循风格参考的结构与深度——但**不要照抄**,要为本档案的主题做合适的调整。",
|
||
"7. 必须包含开头的『一句话浓缩』和结尾的『一句话浓缩』。",
|
||
"8. 必须包含『常见的分析陷阱』和『仍不确定的部分』段落。",
|
||
"",
|
||
"现在开始撰写:",
|
||
])
|
||
|
||
return "\n".join(parts)
|
||
|
||
|
||
# ============================================================
|
||
# API 调用
|
||
# ============================================================
|
||
|
||
def call_api_sync(
|
||
client,
|
||
model: str,
|
||
system_prompt: str,
|
||
user_prompt: str,
|
||
max_tokens: int,
|
||
temperature: float,
|
||
max_retries: int = 3,
|
||
retry_delay: float = 2.0,
|
||
) -> tuple[str, int, int]:
|
||
"""同步调用 API;返回 (content, input_tokens, output_tokens)。带指数退避重试。"""
|
||
last_exception = None
|
||
for attempt in range(1, max_retries + 1):
|
||
try:
|
||
response = client.chat.completions.create(
|
||
model=model,
|
||
messages=[
|
||
{"role": "system", "content": system_prompt},
|
||
{"role": "user", "content": user_prompt},
|
||
],
|
||
max_tokens=max_tokens,
|
||
temperature=temperature,
|
||
)
|
||
content = response.choices[0].message.content or ""
|
||
input_tokens = getattr(response.usage, "prompt_tokens", 0) if response.usage else 0
|
||
output_tokens = getattr(response.usage, "completion_tokens", 0) if response.usage else 0
|
||
return content, input_tokens, output_tokens
|
||
except Exception as e:
|
||
last_exception = e
|
||
if attempt < max_retries:
|
||
wait = retry_delay * (2 ** (attempt - 1))
|
||
logging.warning(f" API 调用失败(第 {attempt} 次):{e};{wait:.1f}s 后重试")
|
||
time.sleep(wait)
|
||
else:
|
||
logging.error(f" API 调用最终失败(已尝试 {max_retries} 次):{e}")
|
||
|
||
raise last_exception # type: ignore
|
||
|
||
|
||
async def call_api_async(
|
||
client,
|
||
model: str,
|
||
system_prompt: str,
|
||
user_prompt: str,
|
||
max_tokens: int,
|
||
temperature: float,
|
||
max_retries: int = 3,
|
||
retry_delay: float = 2.0,
|
||
) -> tuple[str, int, int]:
|
||
"""异步调用 API;带指数退避重试。"""
|
||
last_exception = None
|
||
for attempt in range(1, max_retries + 1):
|
||
try:
|
||
response = await client.chat.completions.create(
|
||
model=model,
|
||
messages=[
|
||
{"role": "system", "content": system_prompt},
|
||
{"role": "user", "content": user_prompt},
|
||
],
|
||
max_tokens=max_tokens,
|
||
temperature=temperature,
|
||
)
|
||
content = response.choices[0].message.content or ""
|
||
input_tokens = getattr(response.usage, "prompt_tokens", 0) if response.usage else 0
|
||
output_tokens = getattr(response.usage, "completion_tokens", 0) if response.usage else 0
|
||
return content, input_tokens, output_tokens
|
||
except Exception as e:
|
||
last_exception = e
|
||
if attempt < max_retries:
|
||
wait = retry_delay * (2 ** (attempt - 1))
|
||
logging.warning(f" API 调用失败(第 {attempt} 次):{e};{wait:.1f}s 后重试")
|
||
await asyncio.sleep(wait)
|
||
else:
|
||
logging.error(f" API 调用最终失败(已尝试 {max_retries} 次):{e}")
|
||
|
||
raise last_exception # type: ignore
|
||
|
||
|
||
# ============================================================
|
||
# 输出清理与验证
|
||
# ============================================================
|
||
|
||
def clean_output(content: str) -> str:
|
||
"""剥离常见的输出包裹问题"""
|
||
content = content.strip()
|
||
|
||
# 剥离首尾的代码块标记
|
||
# 例如 ```markdown\n...\n``` 或 ```\n...\n```
|
||
lines = content.split("\n")
|
||
if lines and lines[0].startswith("```"):
|
||
lines = lines[1:]
|
||
if lines and lines[-1].strip() == "```":
|
||
lines = lines[:-1]
|
||
|
||
content = "\n".join(lines).strip()
|
||
|
||
# 移除常见前言("以下是 ..."、"Here is ...")
|
||
# 仅当第一行不是 Markdown 标题时才考虑
|
||
first_line = content.split("\n", 1)[0] if content else ""
|
||
if first_line and not first_line.startswith("#"):
|
||
prefix_patterns = [
|
||
r"^以下是.*?[::]\s*\n",
|
||
r"^这里是.*?[::]\s*\n",
|
||
r"^Here(?: is| are).*?[::]?\s*\n",
|
||
r"^Below is.*?[::]?\s*\n",
|
||
]
|
||
for pat in prefix_patterns:
|
||
new_content = re.sub(pat, "", content, count=1, flags=re.IGNORECASE | re.DOTALL)
|
||
if new_content != content:
|
||
content = new_content.strip()
|
||
break
|
||
|
||
return content
|
||
|
||
|
||
def validate_output(content: str, item: GenerationItem, master: dict) -> tuple[bool, list[str]]:
|
||
"""基本输出验证;返回 (是否合格, 警告列表)"""
|
||
warnings: list[str] = []
|
||
|
||
# 1. 非空
|
||
if not content or len(content.strip()) < 500:
|
||
return False, ["输出过短或为空"]
|
||
|
||
# 2. 字数检查(粗略)
|
||
length_range = master["type_prompts"].get(item.type, {}).get("length_range", "")
|
||
if length_range:
|
||
# 例如 "8000-15000 字"
|
||
m = re.search(r"(\d+)\s*-\s*(\d+)\s*字", length_range)
|
||
if m:
|
||
min_len, max_len = int(m.group(1)), int(m.group(2))
|
||
actual = len(content)
|
||
# 宽松一点,允许 ±20%
|
||
if actual < min_len * 0.7:
|
||
warnings.append(f"字数 {actual} 显著低于区间 {min_len}-{max_len}")
|
||
elif actual > max_len * 1.3:
|
||
warnings.append(f"字数 {actual} 显著高于区间 {min_len}-{max_len}")
|
||
|
||
# 3. 必含段落检查
|
||
expected_sections = [
|
||
("一句话浓缩", ["一句话浓缩", "一句话定位"]),
|
||
("常见的分析陷阱", ["常见的分析陷阱", "常见陷阱", "分析陷阱"]),
|
||
("仍不确定", ["仍不确定", "不确定的部分"]),
|
||
("与其他文件的关系", ["与其他文件的关系", "相关文件", "关联文件"]),
|
||
]
|
||
for label, candidates in expected_sections:
|
||
found = any(c in content for c in candidates)
|
||
if not found:
|
||
warnings.append(f"缺少段落:{label}")
|
||
|
||
# 4. 是否被代码块包裹(常见错误)
|
||
if content.strip().startswith("```"):
|
||
warnings.append("输出被代码块包裹——可能格式有误")
|
||
|
||
# 任何严重问题(过短)才算失败;其他都是警告
|
||
return True, warnings
|
||
|
||
|
||
# ============================================================
|
||
# 文件生成主流程
|
||
# ============================================================
|
||
|
||
def process_item_sync(
|
||
item: GenerationItem,
|
||
client,
|
||
args: argparse.Namespace,
|
||
master: dict,
|
||
reference_dir: Path,
|
||
output_dir: Path,
|
||
system_prompt: str,
|
||
) -> GenerationResult:
|
||
"""同步处理单一 item"""
|
||
start = time.time()
|
||
|
||
# 检查是否已存在
|
||
output_path = output_dir / item.path
|
||
if output_path.exists() and not args.overwrite:
|
||
logging.info(f" [跳过] 已存在: {item.path}")
|
||
return GenerationResult(
|
||
item=item,
|
||
success=True,
|
||
content=None,
|
||
elapsed_seconds=0.0,
|
||
)
|
||
|
||
# 组装 type prompt
|
||
type_prompt = build_type_prompt(master, item.type)
|
||
full_system = system_prompt + "\n\n" + type_prompt
|
||
|
||
# 加载风格参考
|
||
reference_contents: list[tuple[str, str]] = []
|
||
for ref_path in item.style_reference:
|
||
content = load_reference_file(reference_dir, ref_path)
|
||
if content:
|
||
reference_contents.append((ref_path, content))
|
||
|
||
if not reference_contents:
|
||
logging.warning(f" [警告] 无可用的风格参考文件: {item.id}")
|
||
|
||
# 组装 user prompt
|
||
user_prompt = build_user_prompt(item, reference_contents)
|
||
|
||
# 调用 API
|
||
try:
|
||
content, in_tok, out_tok = call_api_sync(
|
||
client=client,
|
||
model=args.model,
|
||
system_prompt=full_system,
|
||
user_prompt=user_prompt,
|
||
max_tokens=args.max_output_tokens,
|
||
temperature=args.temperature,
|
||
max_retries=args.max_retries,
|
||
retry_delay=args.retry_delay,
|
||
)
|
||
except Exception as e:
|
||
logging.error(f" [失败] API 调用错误: {e}")
|
||
return GenerationResult(
|
||
item=item,
|
||
success=False,
|
||
error=str(e),
|
||
elapsed_seconds=time.time() - start,
|
||
)
|
||
|
||
# 清理输出(剥离代码块、前言等)
|
||
content = clean_output(content)
|
||
|
||
# 验证
|
||
ok, warnings = validate_output(content, item, master)
|
||
if warnings:
|
||
for w in warnings:
|
||
logging.warning(f" [警告] {w}")
|
||
|
||
if not ok:
|
||
# 失败时也存档原始内容,方便诊断
|
||
if content and args.save_failed:
|
||
failed_path = output_dir / "_failed" / f"{item.id}.md"
|
||
failed_path.parent.mkdir(parents=True, exist_ok=True)
|
||
failed_path.write_text(content, encoding="utf-8")
|
||
logging.info(f" 失败的内容已存到: {failed_path}")
|
||
return GenerationResult(
|
||
item=item,
|
||
success=False,
|
||
error="; ".join(warnings),
|
||
content=content,
|
||
elapsed_seconds=time.time() - start,
|
||
input_tokens=in_tok,
|
||
output_tokens=out_tok,
|
||
)
|
||
|
||
# 写出文件
|
||
output_path.parent.mkdir(parents=True, exist_ok=True)
|
||
output_path.write_text(content, encoding="utf-8")
|
||
|
||
elapsed = time.time() - start
|
||
logging.info(
|
||
f" [成功] {item.path} "
|
||
f"({len(content)} 字, {elapsed:.1f}s, in={in_tok}, out={out_tok})"
|
||
)
|
||
|
||
return GenerationResult(
|
||
item=item,
|
||
success=True,
|
||
content=content,
|
||
elapsed_seconds=elapsed,
|
||
input_tokens=in_tok,
|
||
output_tokens=out_tok,
|
||
)
|
||
|
||
|
||
async def process_item_async(
|
||
item: GenerationItem,
|
||
client,
|
||
args: argparse.Namespace,
|
||
master: dict,
|
||
reference_dir: Path,
|
||
output_dir: Path,
|
||
system_prompt: str,
|
||
semaphore: asyncio.Semaphore,
|
||
) -> GenerationResult:
|
||
"""异步处理单一 item(用于并发模式)"""
|
||
async with semaphore:
|
||
start = time.time()
|
||
|
||
output_path = output_dir / item.path
|
||
if output_path.exists() and not args.overwrite:
|
||
logging.info(f" [跳过] 已存在: {item.path}")
|
||
return GenerationResult(item=item, success=True, content=None, elapsed_seconds=0.0)
|
||
|
||
type_prompt = build_type_prompt(master, item.type)
|
||
full_system = system_prompt + "\n\n" + type_prompt
|
||
|
||
reference_contents: list[tuple[str, str]] = []
|
||
for ref_path in item.style_reference:
|
||
content = load_reference_file(reference_dir, ref_path)
|
||
if content:
|
||
reference_contents.append((ref_path, content))
|
||
|
||
user_prompt = build_user_prompt(item, reference_contents)
|
||
|
||
try:
|
||
content, in_tok, out_tok = await call_api_async(
|
||
client=client,
|
||
model=args.model,
|
||
system_prompt=full_system,
|
||
user_prompt=user_prompt,
|
||
max_tokens=args.max_output_tokens,
|
||
temperature=args.temperature,
|
||
max_retries=args.max_retries,
|
||
retry_delay=args.retry_delay,
|
||
)
|
||
except Exception as e:
|
||
logging.error(f" [失败] {item.id}: {e}")
|
||
return GenerationResult(
|
||
item=item, success=False, error=str(e), elapsed_seconds=time.time() - start
|
||
)
|
||
|
||
# 清理输出
|
||
content = clean_output(content)
|
||
|
||
ok, warnings = validate_output(content, item, master)
|
||
if warnings:
|
||
for w in warnings:
|
||
logging.warning(f" [警告] {item.id}: {w}")
|
||
|
||
if not ok:
|
||
if content and args.save_failed:
|
||
failed_path = output_dir / "_failed" / f"{item.id}.md"
|
||
failed_path.parent.mkdir(parents=True, exist_ok=True)
|
||
failed_path.write_text(content, encoding="utf-8")
|
||
logging.info(f" 失败的内容已存到: {failed_path}")
|
||
return GenerationResult(
|
||
item=item, success=False, error="; ".join(warnings),
|
||
content=content, elapsed_seconds=time.time() - start,
|
||
input_tokens=in_tok, output_tokens=out_tok,
|
||
)
|
||
|
||
output_path.parent.mkdir(parents=True, exist_ok=True)
|
||
output_path.write_text(content, encoding="utf-8")
|
||
|
||
elapsed = time.time() - start
|
||
logging.info(
|
||
f" [成功] {item.path} ({len(content)} 字, {elapsed:.1f}s)"
|
||
)
|
||
return GenerationResult(
|
||
item=item, success=True, content=content,
|
||
elapsed_seconds=elapsed, input_tokens=in_tok, output_tokens=out_tok,
|
||
)
|
||
|
||
|
||
# ============================================================
|
||
# 过滤器
|
||
# ============================================================
|
||
|
||
def filter_items(items: list[GenerationItem], args: argparse.Namespace) -> list[GenerationItem]:
|
||
"""根据 CLI 参数过滤 items"""
|
||
filtered = items
|
||
|
||
if args.priority:
|
||
filtered = [x for x in filtered if x.priority == args.priority]
|
||
|
||
if args.topic_group:
|
||
# topic_group 现在是 list,支持多个(OR)
|
||
groups = set(args.topic_group)
|
||
filtered = [x for x in filtered if x.topic_group in groups]
|
||
|
||
if args.type:
|
||
filtered = [x for x in filtered if x.type == args.type]
|
||
|
||
if args.ids:
|
||
id_set = set(args.ids)
|
||
filtered = [x for x in filtered if x.id in id_set]
|
||
|
||
if args.limit and args.limit > 0:
|
||
filtered = filtered[:args.limit]
|
||
|
||
return filtered
|
||
|
||
|
||
# ============================================================
|
||
# 主函数
|
||
# ============================================================
|
||
|
||
def main():
|
||
parser = argparse.ArgumentParser(
|
||
description="台湾 2028 知识库批量生成工具",
|
||
formatter_class=argparse.RawDescriptionHelpFormatter,
|
||
)
|
||
|
||
# API 配置
|
||
parser.add_argument("--api-key", required=True, help="OpenAI 格式 API 的 key")
|
||
parser.add_argument("--base-url", required=True, help="API base URL(如 https://api.openai.com/v1)")
|
||
parser.add_argument("--model", required=True, help="模型名称(如 gpt-4o, deepseek-chat)")
|
||
|
||
# 文件路径
|
||
parser.add_argument("--master-prompt", required=True, help="通用提示词 JSON 文件路径")
|
||
parser.add_argument("--lists", required=True, nargs="+", help="清单 JSON 文件路径(可多个)")
|
||
parser.add_argument("--reference-dir", required=True, help="已写好文件的根目录(用于风格参考)")
|
||
parser.add_argument("--output-dir", required=True, help="生成文件的输出根目录")
|
||
|
||
# 过滤
|
||
parser.add_argument("--priority", help="只跑特定优先级(如 P1, P2, P3, P-MIL-EXT)")
|
||
parser.add_argument("--topic-group", nargs="+",
|
||
help="只跑特定议题组(可多个,OR 关系)。例如军事相关全部:--topic-group taiwan_military military_services weapon_systems operational_concepts threat_assessment")
|
||
parser.add_argument("--type", help="只跑特定类型(如 person, concept)")
|
||
parser.add_argument("--ids", nargs="+", help="只跑特定 ID(可多个)")
|
||
parser.add_argument("--limit", type=int, default=0, help="最多生成几个(0 表示无限制)")
|
||
|
||
# 控制
|
||
parser.add_argument("--dry-run", action="store_true", help="干跑——只列出会跑的 items,不实际调用 API")
|
||
parser.add_argument("--overwrite", action="store_true", help="覆盖已存在的文件")
|
||
parser.add_argument("--concurrency", type=int, default=1, help="并发数(默认 1 顺序执行)")
|
||
parser.add_argument("--max-output-tokens", type=int, default=16000, help="单次输出最大 tokens")
|
||
parser.add_argument("--temperature", type=float, default=0.7, help="模型 temperature")
|
||
parser.add_argument("--max-retries", type=int, default=3, help="API 失败时的最大重试次数")
|
||
parser.add_argument("--retry-delay", type=float, default=2.0, help="重试初始延迟(秒),指数退避")
|
||
parser.add_argument("--save-failed", action="store_true", default=True,
|
||
help="失败的内容也存到 _failed/ 子目录方便诊断(默认开启)")
|
||
parser.add_argument("--no-save-failed", dest="save_failed", action="store_false",
|
||
help="不保存失败的内容")
|
||
parser.add_argument("--summary-file", help="把本次运行的结果汇总写到这个 JSON 文件")
|
||
parser.add_argument("--log-file", help="日志输出到文件(同时也输出到 stdout)")
|
||
parser.add_argument("--verbose", action="store_true", help="详细日志")
|
||
|
||
args = parser.parse_args()
|
||
|
||
# 日志设置
|
||
log_level = logging.DEBUG if args.verbose else logging.INFO
|
||
log_handlers: list[logging.Handler] = [logging.StreamHandler(sys.stdout)]
|
||
if args.log_file:
|
||
log_handlers.append(logging.FileHandler(args.log_file, encoding="utf-8"))
|
||
|
||
logging.basicConfig(
|
||
level=log_level,
|
||
format="%(asctime)s [%(levelname)s] %(message)s",
|
||
handlers=log_handlers,
|
||
force=True,
|
||
)
|
||
|
||
# 路径
|
||
master_path = Path(args.master_prompt)
|
||
list_paths = [Path(p) for p in args.lists]
|
||
reference_dir = Path(args.reference_dir)
|
||
output_dir = Path(args.output_dir)
|
||
|
||
if not reference_dir.exists():
|
||
logging.warning(f"风格参考目录不存在: {reference_dir}(脚本仍会运行,但风格参考会为空)")
|
||
|
||
output_dir.mkdir(parents=True, exist_ok=True)
|
||
|
||
# 加载 master prompt
|
||
logging.info(f"加载通用提示词: {master_path}")
|
||
master = load_master_prompt(master_path)
|
||
|
||
# 加载清单
|
||
logging.info(f"加载清单文件: {[p.name for p in list_paths]}")
|
||
all_items = load_list_files(list_paths)
|
||
logging.info(f"总计 {len(all_items)} 个 items")
|
||
|
||
# 过滤
|
||
items_to_run = filter_items(all_items, args)
|
||
logging.info(f"过滤后待生成 {len(items_to_run)} 个 items")
|
||
|
||
# 干跑模式
|
||
if args.dry_run:
|
||
logging.info("=" * 60)
|
||
logging.info("干跑模式——以下 items 将被生成:")
|
||
logging.info("=" * 60)
|
||
for item in items_to_run:
|
||
output_path = output_dir / item.path
|
||
exists = " [已存在]" if output_path.exists() else ""
|
||
logging.info(f" [{item.priority}] {item.id} → {item.path}{exists}")
|
||
return
|
||
|
||
# 构建 system prompt
|
||
system_prompt = build_system_prompt(master)
|
||
|
||
# 初始化 OpenAI client
|
||
try:
|
||
from openai import OpenAI, AsyncOpenAI
|
||
except ImportError:
|
||
sys.exit("错误: 请安装 openai SDK —— pip install openai")
|
||
|
||
# 顺序执行
|
||
if args.concurrency <= 1:
|
||
client = OpenAI(api_key=args.api_key, base_url=args.base_url)
|
||
results: list[GenerationResult] = []
|
||
for idx, item in enumerate(items_to_run, 1):
|
||
logging.info(f"[{idx}/{len(items_to_run)}] 处理: {item.id} ({item.name})")
|
||
result = process_item_sync(
|
||
item=item,
|
||
client=client,
|
||
args=args,
|
||
master=master,
|
||
reference_dir=reference_dir,
|
||
output_dir=output_dir,
|
||
system_prompt=system_prompt,
|
||
)
|
||
results.append(result)
|
||
|
||
# 并发执行
|
||
else:
|
||
async_client = AsyncOpenAI(api_key=args.api_key, base_url=args.base_url)
|
||
semaphore = asyncio.Semaphore(args.concurrency)
|
||
|
||
async def run_all():
|
||
tasks = [
|
||
process_item_async(
|
||
item=item,
|
||
client=async_client,
|
||
args=args,
|
||
master=master,
|
||
reference_dir=reference_dir,
|
||
output_dir=output_dir,
|
||
system_prompt=system_prompt,
|
||
semaphore=semaphore,
|
||
)
|
||
for item in items_to_run
|
||
]
|
||
return await asyncio.gather(*tasks, return_exceptions=False)
|
||
|
||
logging.info(f"并发模式:并发数 {args.concurrency}")
|
||
results = asyncio.run(run_all())
|
||
|
||
# 汇总
|
||
succeeded = [r for r in results if r.success]
|
||
failed = [r for r in results if not r.success]
|
||
total_input = sum(r.input_tokens for r in results)
|
||
total_output = sum(r.output_tokens for r in results)
|
||
total_time = sum(r.elapsed_seconds for r in results)
|
||
|
||
logging.info("=" * 60)
|
||
logging.info(f"完成: 总计 {len(results)} | 成功 {len(succeeded)} | 失败 {len(failed)}")
|
||
logging.info(f"总耗时: {total_time:.1f}s ({total_time / 60:.1f} 分钟)")
|
||
logging.info(f"Tokens: input {total_input:,} | output {total_output:,}")
|
||
|
||
if failed:
|
||
logging.info("失败列表:")
|
||
for r in failed:
|
||
logging.info(f" - {r.item.id} ({r.item.path}): {r.error}")
|
||
|
||
logging.info("=" * 60)
|
||
|
||
# 写汇总 JSON
|
||
if args.summary_file:
|
||
summary = {
|
||
"run_at": time.strftime("%Y-%m-%d %H:%M:%S"),
|
||
"model": args.model,
|
||
"base_url": args.base_url,
|
||
"total": len(results),
|
||
"succeeded": len(succeeded),
|
||
"failed": len(failed),
|
||
"total_elapsed_seconds": total_time,
|
||
"total_input_tokens": total_input,
|
||
"total_output_tokens": total_output,
|
||
"items": [
|
||
{
|
||
"id": r.item.id,
|
||
"path": r.item.path,
|
||
"name": r.item.name,
|
||
"success": r.success,
|
||
"error": r.error,
|
||
"elapsed_seconds": round(r.elapsed_seconds, 2),
|
||
"input_tokens": r.input_tokens,
|
||
"output_tokens": r.output_tokens,
|
||
"content_length": len(r.content) if r.content else 0,
|
||
}
|
||
for r in results
|
||
],
|
||
"failed_ids": [r.item.id for r in failed],
|
||
}
|
||
summary_path = Path(args.summary_file)
|
||
summary_path.parent.mkdir(parents=True, exist_ok=True)
|
||
summary_path.write_text(
|
||
json.dumps(summary, ensure_ascii=False, indent=2),
|
||
encoding="utf-8",
|
||
)
|
||
logging.info(f"汇总已写到: {summary_path}")
|
||
|
||
if failed:
|
||
logging.info(
|
||
f"提示:重跑失败的 items,用 --ids {' '.join(r.item.id for r in failed[:5])}..."
|
||
)
|
||
|
||
|
||
if __name__ == "__main__":
|
||
main()
|