Files
human-voice-rewrite-demo/main.py
T

204 lines
7.4 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""Human Voice Rewrite Demo backend (PRD v1.1).
Serves the static demo page and domain endpoints that run the real
OpenRouter pipeline. Run via ./local_start.sh.
"""
from __future__ import annotations
from pathlib import Path
from typing import Any, Callable
from fastapi import FastAPI, HTTPException
from fastapi.responses import FileResponse
from evidence import sanitize_diagnosis
from llm import LLMError, LlmClient
from prompts import (
build_analyze_prompt,
build_diagnose_prompt,
build_recheck_prompt,
build_reference_prompt,
build_scaffold_prompt,
normalize_diagnose_payload,
)
from schemas import (
AnalyzeRequest,
AnalyzeResponse,
DiagnoseRequest,
DiagnoseResponse,
GlobalChecks,
RecheckRequest,
RecheckResponse,
ReferenceRequest,
ReferenceResponse,
RevisionTarget,
ScaffoldRequest,
ScaffoldResponse,
)
app = FastAPI(title="Human Voice Rewrite Demo", version="1.1.0")
client: LlmClient | None = None
HERE = Path(__file__).resolve().parent
def get_client() -> LlmClient:
global client
if client is None:
client = LlmClient()
return client
def run_json(phase: str, system: str, user: str, check: Callable[[dict[str, Any]], str | None] | None = None) -> dict:
try:
return get_client().complete_json(system, user, validate=check)
except LLMError as exc:
raise HTTPException(status_code=502, detail=f"{phase}失败:{exc}") from exc
def _count_check(expected: int, key: str) -> Callable[[dict[str, Any]], str | None]:
"""返回语义校验器:LLM 输出里按段组织的数组条目数必须与输入段数一致。
不一致时返回一行中文说明(触发 complete_json 重试一次;两次都失败才变
502 单行错误)——真因:真实 LLM 偶发返回合法 JSON 但段落条目不全,
此前无数量校验被静默接受,AI 初审出现"4 段标题只显示 1 段"(用户实测)。"""
def check(data: dict[str, Any]) -> str | None:
got = len(data.get(key) or [])
if got != expected:
return f"结果段落数与输入不一致(输入 {expected} 段,返回 {got} 段)"
return None
return check
def _friendly_validation(exc: Exception) -> str:
"""把 pydantic ValidationError 压缩成一行用户可读摘要(PRD §11 失败兜底:
返回用户可读错误,不把完整校验堆栈甩给页面)。"""
errs = getattr(exc, "errors", lambda: [])()
if not errs:
return f"结果格式不符合预期:{exc}"
parts = []
for e in errs[:2]:
loc = " → ".join(str(x) for x in e.get("loc", []))
parts.append(f"{loc}{e.get('msg', '')}" if loc else e.get("msg", ""))
more = f"(共 {len(errs)} 处)" if len(errs) > 2 else ""
return "结果格式不符合预期:" + "".join(parts) + more
@app.get("/api/health")
def health() -> dict:
return {"status": "ok"}
@app.post("/api/analyze", response_model=AnalyzeResponse)
def analyze(req: AnalyzeRequest) -> AnalyzeResponse:
if not any(p.strip() for p in req.paragraphs):
raise HTTPException(status_code=422, detail="文本为空,无法开始分析")
system, user = build_analyze_prompt(req.prompt, req.word_limit, req.paragraphs, req.constraints)
data = run_json("分析", system, user, check=_count_check(len(req.paragraphs), "paragraphs"))
try:
return AnalyzeResponse.model_validate(data)
except Exception as exc: # pydantic ValidationError etc.
raise HTTPException(status_code=502, detail=f"分析{_friendly_validation(exc)}") from exc
@app.post("/api/diagnose", response_model=DiagnoseResponse)
def diagnose(req: DiagnoseRequest) -> DiagnoseResponse:
system, user = build_diagnose_prompt(
req.prompt,
req.word_limit,
req.paragraphs,
req.confirmed_anchors,
req.global_constraints,
req.paragraph_constraints,
req.initial_analysis,
)
data = run_json("诊断", system, user, check=_count_check(len(req.paragraphs), "paragraph_briefs"))
data = normalize_diagnose_payload(data)
try:
resp = DiagnoseResponse.model_validate(data)
except Exception as exc:
raise HTTPException(status_code=502, detail=f"诊断{_friendly_validation(exc)}") from exc
return sanitize_diagnosis(resp, req.paragraphs)
@app.post("/api/recheck", response_model=RecheckResponse)
def recheck(req: RecheckRequest) -> RecheckResponse:
if len(req.original_paragraphs) != len(req.rewrite_paragraphs):
raise HTTPException(status_code=422, detail="原文段数与改写段数不一致")
if any(not p.strip() for p in req.rewrite_paragraphs):
raise HTTPException(status_code=422, detail="存在空段:全部段落完成改写后才能提交复检")
system, user = build_recheck_prompt(req)
data = run_json("复检", system, user)
if isinstance(data.get("status"), str):
data["status"] = "pass" if data["status"] == "pass" else "revision_required"
if not data.get("checked_rewrite_version"):
data["checked_rewrite_version"] = req.rewrite_version or ""
try:
resp = RecheckResponse.model_validate(data)
except Exception as exc:
raise HTTPException(status_code=502, detail=f"复检{_friendly_validation(exc)}") from exc
if resp.status == "revision_required" and not resp.revision_targets:
resp = RecheckResponse(
status="revision_required",
checked_rewrite_version=resp.checked_rewrite_version or req.rewrite_version,
global_checks=resp.global_checks or GlobalChecks(),
revision_targets=[
RevisionTarget(
paragraph_id="p1",
blocking_issue="模型判定需要返工但未指明段落",
single_revision_goal="重读当前段改写并修正最明显的问题",
)
],
)
return resp
@app.post("/api/scaffold", response_model=ScaffoldResponse)
def scaffold(req: ScaffoldRequest) -> ScaffoldResponse:
system, user = build_scaffold_prompt(
req.paragraph_id,
req.original_text,
req.semantic_anchor,
req.rewrite_goal,
req.global_constraints,
req.paragraph_constraints,
)
data = run_json("写作起点", system, user)
try:
resp = ScaffoldResponse.model_validate(data)
except Exception as exc:
raise HTTPException(status_code=502, detail=f"写作起点{_friendly_validation(exc)}") from exc
resp.paragraph_id = resp.paragraph_id or req.paragraph_id
return resp
@app.post("/api/reference", response_model=ReferenceResponse)
def reference(req: ReferenceRequest) -> ReferenceResponse:
goal = req.rewrite_goal or req.primary_goal
system, user = build_reference_prompt(
req.paragraph_id,
req.original_text,
req.semantic_anchor,
goal,
req.global_constraints,
req.paragraph_constraints,
)
data = run_json("参考片段", system, user)
try:
resp = ReferenceResponse.model_validate(data)
except Exception as exc:
raise HTTPException(status_code=502, detail=f"参考片段{_friendly_validation(exc)}") from exc
snippet = resp.reference_snippet or resp.starter
resp.paragraph_id = resp.paragraph_id or req.paragraph_id
resp.starter = snippet
resp.reference_snippet = snippet
return resp
@app.get("/")
def index() -> FileResponse:
return FileResponse(HERE / "index.html", media_type="text/html")