192 lines
6.5 KiB
Python
192 lines
6.5 KiB
Python
# IR 生成、验证、Diff 接口
|
|
|
|
import asyncio
|
|
import json
|
|
import logging
|
|
from concurrent.futures import ThreadPoolExecutor
|
|
|
|
from fastapi import APIRouter, HTTPException, Query
|
|
from fastapi.responses import StreamingResponse
|
|
|
|
from server.services.ir_engine.generator import ir_generator
|
|
from server.services.ir_engine.pipeline import ir_pipeline
|
|
from server.services.ir_engine.validator import ir_validator
|
|
from server.services.ir_engine.diff import ir_diff
|
|
from server.services.prd_manager.service import prd_service
|
|
|
|
logger = logging.getLogger(__name__)
|
|
router = APIRouter()
|
|
|
|
_pipeline_executor = ThreadPoolExecutor(max_workers=2)
|
|
|
|
|
|
@router.get("/generate-stream")
|
|
async def generate_ir_stream(prd_id: str = Query(...), skill_name: str = "default"):
|
|
"""基于 PRD 生成 IR,通过 SSE 流式返回进度"""
|
|
prd = await prd_service.get_prd(prd_id)
|
|
if not prd:
|
|
raise HTTPException(404, "PRD not found")
|
|
|
|
async def event_stream():
|
|
# Wait for background parsing if still in progress
|
|
if prd.get("status") == "parsing":
|
|
yield f"data: {json.dumps({'stage': 0, 'status': 'running', 'message': '文档解析中,请稍候...', 'detail': '图片分析可能需要 30-60 秒', 'stage_total': 4})}\n\n"
|
|
waited = 0
|
|
while waited < 120:
|
|
await asyncio.sleep(1)
|
|
waited += 1
|
|
p = await prd_service.get_prd(prd_id)
|
|
if not p or p.get("status") in ("ready", "error"):
|
|
break
|
|
if waited % 5 == 0:
|
|
yield f"data: {json.dumps({'stage': 0, 'status': 'running', 'message': f'文档解析中...({waited}s)', 'stage_total': 4})}\n\n"
|
|
p = await prd_service.get_prd(prd_id)
|
|
if p and p.get("status") == "error":
|
|
yield f"data: {json.dumps({'error': p.get('error', '解析失败'), 'stage_total': 4})}\n\n"
|
|
return
|
|
|
|
# Reload after parsing
|
|
prd2 = await prd_service.get_prd(prd_id)
|
|
parsed_doc = _load_parsed_doc(prd2 or prd)
|
|
if not parsed_doc:
|
|
yield f"data: {json.dumps({'error': '无法加载解析后的文档'})}\n\n"
|
|
return
|
|
|
|
yield f"data: {json.dumps({'stage': 0, 'status': 'done', 'message': '文档解析完成', 'stage_total': 4})}\n\n"
|
|
|
|
# Run pipeline in thread → Queue → SSE
|
|
queue: asyncio.Queue = asyncio.Queue()
|
|
|
|
def _run_in_thread():
|
|
try:
|
|
for event in ir_pipeline.run_streaming(prd_id, parsed_doc):
|
|
queue.put_nowait(event)
|
|
except Exception as e:
|
|
logger.exception("Pipeline error")
|
|
queue.put_nowait({"error": str(e), "done": True})
|
|
finally:
|
|
queue.put_nowait(None)
|
|
|
|
loop = asyncio.get_event_loop()
|
|
loop.run_in_executor(_pipeline_executor, _run_in_thread)
|
|
|
|
while True:
|
|
event = await queue.get()
|
|
if event is None:
|
|
break
|
|
yield f"data: {json.dumps(event, ensure_ascii=False)}\n\n"
|
|
yield "data: {\"done\": true}\n\n"
|
|
|
|
return StreamingResponse(
|
|
event_stream(),
|
|
media_type="text/event-stream",
|
|
headers={
|
|
"Cache-Control": "no-cache",
|
|
"Connection": "keep-alive",
|
|
"X-Accel-Buffering": "no",
|
|
},
|
|
)
|
|
|
|
|
|
def _load_parsed_doc(prd: dict) -> dict | None:
|
|
"""Load parsed document from disk or build from PRD data."""
|
|
import os as _os
|
|
parsed_path = prd.get("parsed_path", "")
|
|
if parsed_path and _os.path.isfile(parsed_path):
|
|
try:
|
|
with open(parsed_path, "r", encoding="utf-8") as f:
|
|
return json.load(f)
|
|
except Exception:
|
|
pass
|
|
|
|
sections = prd.get("sections", [])
|
|
if not sections and prd.get("full_text"):
|
|
sections = [{"source": "正文", "blocks": [{"type": "para", "index": 1, "text": prd["full_text"]}], "images": []}]
|
|
|
|
if sections or prd.get("full_text"):
|
|
return {
|
|
"source": "",
|
|
"sections": sections,
|
|
"image_sources": {},
|
|
"image_analysis": prd.get("images", []),
|
|
"resolved_conflicts": [],
|
|
}
|
|
return None
|
|
|
|
|
|
@router.post("/generate")
|
|
async def generate_ir(prd_id: str, skill_name: str = "default"):
|
|
"""基于 PRD 和 Skill 生成 IR"""
|
|
prd = await prd_service.get_prd(prd_id)
|
|
if not prd:
|
|
raise HTTPException(404, "PRD not found")
|
|
|
|
prd_text = prd.get("full_text", "")
|
|
if not prd_text:
|
|
raise HTTPException(400, "PRD has no text content")
|
|
|
|
result = await ir_generator.generate(prd_id, prd_text, skill_name)
|
|
if "error" in result:
|
|
raise HTTPException(500, result["error"])
|
|
|
|
return {
|
|
"ir_id": result["ir_id"],
|
|
"prd_id": result["prd_id"],
|
|
"yaml_content": result["yaml_content"],
|
|
"ir_json": result.get("ir_json", {}),
|
|
"audit": result.get("audit", {}),
|
|
"audit_report": result.get("audit_report", ""),
|
|
"skill_used": result["skill_used"],
|
|
"created_at": result["created_at"],
|
|
"pipeline_stats": result.get("pipeline_stats", {}),
|
|
}
|
|
|
|
|
|
@router.get("/{ir_id}")
|
|
async def get_ir(ir_id: str):
|
|
"""获取 IR 详情"""
|
|
ir = await ir_generator.get_ir(ir_id)
|
|
if not ir:
|
|
raise HTTPException(404, "IR not found")
|
|
return ir
|
|
|
|
|
|
@router.put("/{ir_id}")
|
|
async def update_ir(ir_id: str, body: dict):
|
|
"""更新 IR 内容(编辑后同步回后端)"""
|
|
ir = await ir_generator.get_ir(ir_id)
|
|
if not ir:
|
|
raise HTTPException(404, "IR not found")
|
|
new_content = body.get("yaml_content", "")
|
|
if new_content:
|
|
ir["yaml_content"] = new_content
|
|
try:
|
|
ir["ir_json"] = json.loads(new_content)
|
|
except json.JSONDecodeError:
|
|
pass
|
|
return {"ir_id": ir_id, "status": "updated"}
|
|
|
|
|
|
@router.post("/{ir_id}/validate")
|
|
async def validate_ir(ir_id: str):
|
|
"""验证 IR 是否符合 Schema 和 Principles"""
|
|
ir = await ir_generator.get_ir(ir_id)
|
|
if not ir:
|
|
raise HTTPException(404, "IR not found")
|
|
|
|
result = await ir_validator.validate(ir["yaml_content"], ir.get("skill_used", "default"))
|
|
return {
|
|
"ir_id": ir_id,
|
|
"valid": result["valid"],
|
|
"issues": result["issues"],
|
|
}
|
|
|
|
|
|
@router.get("/diff")
|
|
async def diff_ir(ir_id_a: str, ir_id_b: str):
|
|
"""对比两个 IR 版本的差异"""
|
|
result = await ir_diff.diff(ir_id_a, ir_id_b)
|
|
if "error" in result:
|
|
raise HTTPException(400, result["error"])
|
|
return result
|