201 lines
7.1 KiB
Python
Executable File
201 lines
7.1 KiB
Python
Executable File
"""
|
|
Stage 1: Semantic Index Generation.
|
|
|
|
Reads the full parsed PRD JSON, calls the LLM once to produce a semantic index
|
|
identifying all function units, concepts, and their document sources.
|
|
|
|
Output: output/semantic_index.json
|
|
"""
|
|
|
|
import json
|
|
import re
|
|
import sys
|
|
import time
|
|
from pathlib import Path
|
|
|
|
import config
|
|
|
|
|
|
def format_document_for_prompt(doc: dict) -> str:
|
|
"""Render the full parsed document as a readable string for the LLM prompt."""
|
|
lines = []
|
|
|
|
# ---- Sections ----
|
|
lines.append("=== SECTIONS ===")
|
|
for i, section in enumerate(doc.get("sections", [])):
|
|
source = section.get("source", f"(无标题-章节{i})")
|
|
lines.append(f"\n--- Section: {source} ---")
|
|
|
|
for block in section.get("blocks", []):
|
|
if block["type"] == "para":
|
|
lines.append(f"[段落 {block['index']}] {block['text']}")
|
|
elif block["type"] == "table":
|
|
lines.append(f"[表格 {block.get('table', '?')}]")
|
|
headers = block.get("headers", [])
|
|
lines.append(f" 表头: {' | '.join(headers)}")
|
|
for row in block.get("rows", []):
|
|
cols = row.get("columns", [])
|
|
cell_texts = []
|
|
for c in cols:
|
|
cell_texts.append(f"[行{c.get('row','?')}]{c.get('name','')}: {c.get('text','')}")
|
|
lines.append(f" {'; '.join(cell_texts)}")
|
|
|
|
images = section.get("images", [])
|
|
if images:
|
|
lines.append(f" 图片引用: {', '.join(images)}")
|
|
|
|
# ---- Image Analysis ----
|
|
lines.append("\n\n=== IMAGE_ANALYSIS (流程图逻辑树) ===")
|
|
for img in doc.get("image_analysis", []):
|
|
rid = img.get("rid", "?")
|
|
img_type = img.get("type", "?")
|
|
lines.append(f"\n--- Image: {rid} (type={img_type}) ---")
|
|
lines.append(f" 描述: {img.get('description', '')[:300]}")
|
|
|
|
lt = img.get("logic_tree")
|
|
if lt:
|
|
lines.append(f" 逻辑树根节点: {lt.get('root', '?')}")
|
|
lines.append(" 节点详情:")
|
|
for node in lt.get("nodes", []):
|
|
nid = node.get("id", "?")
|
|
ntype = node.get("type", "?")
|
|
desc = node.get("description", "") or node.get("condition", "")
|
|
lines.append(f" [{ntype}] {nid}: {desc}")
|
|
branches = node.get("branches", [])
|
|
if branches:
|
|
for br in branches:
|
|
lines.append(f" → {br['value']} → {br['target']}")
|
|
|
|
# ---- Resolved Conflicts ----
|
|
conflicts = doc.get("resolved_conflicts", [])
|
|
if conflicts:
|
|
lines.append("\n\n=== RESOLVED_CONFLICTS (图文冲突仲裁) ===")
|
|
for c in conflicts:
|
|
lines.append(
|
|
f" [{c.get('conflict_type','?')}] {c.get('section','?')}: "
|
|
f"以{c.get('source','?')}为准 — {c.get('correction','')}"
|
|
)
|
|
|
|
return "\n".join(lines)
|
|
|
|
|
|
def build_prompt(doc: dict) -> str:
|
|
"""Load the prompt template and inject the formatted document."""
|
|
template_path = Path(config.PROMPTS_DIR) / "step1_semantic_index.txt"
|
|
template = template_path.read_text(encoding="utf-8")
|
|
|
|
formatted_doc = format_document_for_prompt(doc)
|
|
prompt = template.replace("{document_json}", formatted_doc)
|
|
return prompt
|
|
|
|
|
|
def extract_json_from_response(text: str) -> str:
|
|
"""Robustly extract JSON from LLM response, handling markdown fences."""
|
|
# Try ```json ... ``` first
|
|
m = re.search(r"```(?:json)?\s*([\s\S]*?)```", text)
|
|
if m:
|
|
return m.group(1).strip()
|
|
|
|
# Try to find the outermost { ... }
|
|
start = text.find("{")
|
|
if start == -1:
|
|
raise ValueError("No JSON object found in LLM response")
|
|
|
|
depth = 0
|
|
for i in range(start, len(text)):
|
|
if text[i] == "{":
|
|
depth += 1
|
|
elif text[i] == "}":
|
|
depth -= 1
|
|
if depth == 0:
|
|
return text[start : i + 1]
|
|
|
|
raise ValueError("Unclosed JSON object in LLM response")
|
|
|
|
|
|
def call_llm(prompt: str, max_retries: int = 2) -> dict:
|
|
"""Send the prompt to the LLM and return the parsed semantic index JSON."""
|
|
client = config.llm_client()
|
|
|
|
for attempt in range(max_retries + 1):
|
|
print(f" LLM 调用 (尝试 {attempt + 1}/{max_retries + 1})...", flush=True)
|
|
try:
|
|
resp = client.chat.completions.create(
|
|
model=config.MODEL_NAME,
|
|
messages=[
|
|
{"role": "system", "content": "你是一个精确的 JSON 输出引擎。只输出合法的 JSON,不输出任何其他文字。"},
|
|
{"role": "user", "content": prompt},
|
|
],
|
|
temperature=config.TEMPERATURE,
|
|
max_tokens=config.MAX_TOKENS,
|
|
)
|
|
content = resp.choices[0].message.content
|
|
if content is None:
|
|
raise RuntimeError("LLM returned empty response")
|
|
|
|
json_str = extract_json_from_response(content)
|
|
result = json.loads(json_str)
|
|
|
|
# Quick structural validation
|
|
_validate_schema(result)
|
|
return result
|
|
|
|
except (json.JSONDecodeError, ValueError) as e:
|
|
print(f" JSON 解析失败: {e}")
|
|
if attempt < max_retries:
|
|
time.sleep(2)
|
|
else:
|
|
raise RuntimeError(f"无法从 LLM 响应中解析 JSON: {e}") from e
|
|
|
|
|
|
def _validate_schema(data: dict) -> None:
|
|
"""Validate the top-level schema of the semantic index."""
|
|
if "feature_name" not in data:
|
|
raise ValueError("semantic_index 缺少 'feature_name' 字段")
|
|
if "function_units" not in data:
|
|
raise ValueError("semantic_index 缺少 'function_units' 字段")
|
|
|
|
units = data["function_units"]
|
|
if not isinstance(units, list) or len(units) == 0:
|
|
raise ValueError("function_units 必须是非空数组")
|
|
|
|
for i, fu in enumerate(units):
|
|
if not fu.get("unit_id"):
|
|
raise ValueError(f"function_unit[{i}] 缺少 unit_id")
|
|
if not fu.get("name"):
|
|
raise ValueError(f"function_unit[{i}] 缺少 name")
|
|
|
|
|
|
def main():
|
|
print("=" * 60)
|
|
print("阶段一:宏观语义索引 (Semantic Index)")
|
|
print("=" * 60)
|
|
|
|
# 1. Load input
|
|
print(f"\n[1/4] 加载输入文档: {config.INPUT_JSON}")
|
|
doc = config.load_input_document()
|
|
print(f" 已加载 {len(doc.get('sections', []))} 个 section, "
|
|
f"{len(doc.get('image_analysis', []))} 张图片分析")
|
|
|
|
# 2. Build prompt
|
|
print(f"\n[2/4] 构造 Prompt...")
|
|
prompt = build_prompt(doc)
|
|
print(f" Prompt 长度: {len(prompt)} 字符 (~{len(prompt)//3} tokens)")
|
|
|
|
# 3. Call LLM
|
|
print(f"\n[3/4] 调用 LLM ({config.MODEL_NAME})...")
|
|
semantic_index = call_llm(prompt)
|
|
|
|
# 4. Save result
|
|
print(f"\n[4/4] 保存语义索引: {config.SEMANTIC_INDEX_JSON}")
|
|
config.save_json(semantic_index, config.SEMANTIC_INDEX_JSON)
|
|
|
|
n_concepts = len(semantic_index.get("concepts", []))
|
|
n_units = len(semantic_index.get("function_units", []))
|
|
print(f"\n完成! 提取了 {n_concepts} 个概念, {n_units} 个功能单元.")
|
|
print(f"输出: {config.SEMANTIC_INDEX_JSON}")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|