Files
2026Technology-Competition/src/genesis/server/app.py
T
lhl 80daadcd31 fix(websocket): 修复 app.py 不可导入并新增真实链路验证
- 将 /api/sessions/{sid}/ws 端点移入 create_app(此前置于模块级导致整模块 import NameError,回归被验证拦截)
- register_loop + subscribe 调整至 accept 之前,缩小连接已开但未订阅期间的进度丢失窗口
- 新增 tests/test_verify_ws_real_flow.py:驱动真实 HTTP 聊天流程断言 WS 收到 agent 实际发射的 parse/impact 进度
- 同步 WebSocket 计划文档 Task 3 代码片段(标注端点必须位于 create_app 内)
- 全量 pytest 实测 583 passed / 99.03% 达标
2026-08-29 14:36:12 +08:00

355 lines
13 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.
"""Web 服务化:FastAPI 端点(S4)。
实现 api-design.md §2 的核心端点(v1 同步执行 + 内嵌零构建前端)。
- 会话:POST/GET /api/sessions、GET /api/sessions/{id}、DELETE
- 文件:POST /api/sessions/{id}/filesmultipart
- 解析:POST start-parse、GET parse-result、POST confirm-parse
- 影响:POST start-impact、GET impact-result、POST confirm-impact
- 生成/QAPOST generate、POST run-qa、GET qa-result
- 结果:GET result/preview、result/download、result/impact-report、result/qa-report
- 前端:GET / 返回内嵌单页
"""
from __future__ import annotations
import asyncio
from pathlib import Path
from typing import Any
from fastapi import FastAPI, File, Form, HTTPException, UploadFile, WebSocket, WebSocketDisconnect
from fastapi.responses import FileResponse, HTMLResponse
from pydantic import BaseModel
from genesis.inference.factory import build_inference_engine
from genesis.server.hub import hub
from genesis.server.service import FileTypeError, GenesisService, ServiceStepError
from genesis.server.store import (
ProjectConfigError, ProjectsStore, SessionNotFoundError, SessionStore,
)
VERSION = "0.1.0"
class _FakeEngine:
"""离线确定性引擎(--engine fake):用于无 API key 的 Web 演示/测试。"""
def chat_structured(self, *, session_id, prompt, variables, schema, retry_count=2):
from types import SimpleNamespace
title = variables.get("title", "x")
return SimpleNamespace(
data={
"title": title,
"blocks": [{"type": "paragraph",
"text": "本機能はFakeLLMにより生成された十分な説明内容であり、書込規則を満たす。"}],
},
status="ok",
)
class ChatMessageReq(BaseModel):
content: str
class SessionCreate(BaseModel):
user_id: str = "default"
name: str | None = None
project: str | None = None
class ProjectCreate(BaseModel):
name: str
display_name: str = ""
template: str = ""
write_instruction: str = ""
rules: list[str] = []
existing_system_code_dir: str = ""
design_docs_dir: str = ""
class FileUploadResp(BaseModel):
file_id: str
file_name: str
size: int
class GenerateReq(BaseModel):
output_language: str = "auto"
def _error(status: int, code: str, message: str) -> HTTPException:
return HTTPException(status_code=status, detail={"code": code, "message": message})
def create_app(
store: SessionStore | None = None,
data_root: str = "data/server",
engine: Any = None,
) -> FastAPI:
store = store or SessionStore()
is_fake = engine == "fake"
if engine == "fake":
engine = _FakeEngine()
elif engine is None:
engine = None # 真实模式:generate/qa 时按需 build(避免未配置 key 直接 503)
projects = ProjectsStore(db_path=str(store._db)) if store else None
service = GenesisService(store=store, data_root=data_root, engine=engine, projects=projects)
from genesis.chat.agent import ChatAgent
chat_agent = ChatAgent(service=service, fake=is_fake, engine=engine)
static_dir = Path(__file__).parent / "static"
chat_html = (static_dir / "chat.html").read_text(encoding="utf-8") if (static_dir / "chat.html").exists() else "<html><body>Genesis Chat</body></html>"
app = FastAPI(title="Genesis API", version=VERSION)
# ---------- 健康/前端 ----------
@app.get("/api/health")
def health():
return {"status": "ok", "version": VERSION}
@app.get("/", response_class=HTMLResponse)
def index():
return chat_html
@app.get("/chat_state.js")
def chat_state_js():
p = static_dir / "chat_state.js"
if not p.exists():
raise _error(404, "NOT_FOUND", "chat_state.js 不存在")
return FileResponse(p, media_type="application/javascript")
@app.get("/chat_ws.js")
def chat_ws_js():
p = static_dir / "chat_ws.js"
if not p.exists():
raise _error(404, "NOT_FOUND", "chat_ws.js 不存在")
return FileResponse(p, media_type="application/javascript")
# ---------- 项目配置 ----------
@app.post("/api/projects")
def create_project(body: ProjectCreate):
try:
cfg = service.projects.upsert(
name=body.name, display_name=body.display_name, template=body.template,
write_instruction=body.write_instruction, rules=body.rules,
existing_system_code_dir=body.existing_system_code_dir,
design_docs_dir=body.design_docs_dir,
)
except ProjectConfigError as e:
raise _error(400, "PROJECT_CONFIG_INVALID", str(e))
return cfg.to_dict
@app.get("/api/projects")
def list_projects():
if not service.projects:
return []
return [c.to_dict for c in service.projects.list()]
@app.get("/api/projects/{name}")
def get_project(name: str):
if not service.projects:
raise _error(404, "PROJECT_NOT_FOUND", f"项目不存在: {name}")
cfg = service.projects.get(name)
if cfg is None:
raise _error(404, "PROJECT_NOT_FOUND", f"项目不存在: {name}")
return cfg.to_dict
@app.delete("/api/projects/{name}")
def delete_project(name: str):
if not service.projects:
return {"deleted": False}
return {"deleted": service.projects.delete(name)}
# ---------- 会话 ----------
@app.post("/api/sessions")
def create_session(body: SessionCreate):
rec = service.create_session(body.user_id, name=body.name, project=body.project)
return {"session_id": rec.session_id, "status": rec.status, "name": rec.name}
@app.get("/api/sessions")
def list_sessions(user_id: str = "default", project: str | None = None):
return [
{"session_id": r.session_id, "name": r.name, "project": r.project,
"status": r.status, "updated_at": r.updated_at}
for r in service.store.list_sessions(user_id, project=project)
]
@app.get("/api/sessions/{sid}")
def get_session(sid: str):
try:
rec = service.get_session(sid)
except SessionNotFoundError:
raise _error(404, "SESSION_NOT_FOUND", f"会话不存在: {sid}")
return rec.to_dict
@app.delete("/api/sessions/{sid}")
def delete_session(sid: str):
ok = service.store.delete_session(sid)
return {"deleted": ok}
# ---------- 文件上传 ----------
@app.post("/api/sessions/{sid}/files", response_model=FileUploadResp)
async def upload_file(sid: str, file_type: str = Form(...), file: UploadFile = File(...)):
content = await file.read()
try:
entry = service.upload_file(sid, file_type, file.filename or "upload", content)
except FileTypeError as e:
raise _error(400, "FILE_TYPE_INVALID", str(e))
except SessionNotFoundError:
raise _error(404, "SESSION_NOT_FOUND", f"会话不存在: {sid}")
return FileUploadResp(file_id=entry["file_id"], file_name=entry["name"], size=entry["size"])
# ---------- 解析 ----------
@app.post("/api/sessions/{sid}/start-parse")
def start_parse(sid: str):
try:
rec = service.run_parse(sid)
except ServiceStepError as e:
raise _error(409, "STATE_TRANSITION_INVALID", str(e))
except SessionNotFoundError:
raise _error(404, "SESSION_NOT_FOUND", f"会话不存在: {sid}")
return {"ok": True, "status": rec.status}
@app.get("/api/sessions/{sid}/parse-result")
def parse_result(sid: str):
rec = service.get_session(sid)
import json
return json.loads(rec.structured_summary) if rec.structured_summary else {}
@app.post("/api/sessions/{sid}/confirm-parse")
def confirm_parse(sid: str):
try:
rec = service.confirm_parse(sid)
except ServiceStepError as e:
raise _error(409, "STATE_TRANSITION_INVALID", str(e))
return {"ok": True, "status": rec.status}
# ---------- 影响调查 ----------
@app.post("/api/sessions/{sid}/start-impact")
def start_impact(sid: str):
try:
rec = service.run_impact(sid)
except ServiceStepError as e:
raise _error(409, "STATE_TRANSITION_INVALID", str(e))
return {"ok": True, "status": rec.status}
@app.get("/api/sessions/{sid}/impact-result")
def impact_result(sid: str):
rec = service.get_session(sid)
import json
return json.loads(rec.impact_summary) if rec.impact_summary else {}
@app.post("/api/sessions/{sid}/confirm-impact")
def confirm_impact(sid: str):
try:
rec = service.confirm_impact(sid)
except ServiceStepError as e:
raise _error(409, "STATE_TRANSITION_INVALID", str(e))
return {"ok": True, "status": rec.status}
# ---------- 生成 / QA ----------
@app.post("/api/sessions/{sid}/generate")
def generate(sid: str, body: GenerateReq | None = None):
lang = (body.output_language if body else "auto") or "auto"
try:
if service.engine is None:
service.engine = build_inference_engine()
rec = service.run_generate(sid, output_language=lang)
except ServiceStepError as e:
raise _error(409, "STATE_TRANSITION_INVALID", str(e))
return {"ok": True, "status": rec.status, "result_path": rec.result_path}
@app.post("/api/sessions/{sid}/run-qa")
def run_qa(sid: str):
try:
rec = service.run_qa(sid)
except ServiceStepError as e:
raise _error(409, "STATE_TRANSITION_INVALID", str(e))
return {"ok": True, "status": rec.status}
@app.get("/api/sessions/{sid}/qa-result")
def qa_result(sid: str):
rec = service.get_session(sid)
import json
return json.loads(rec.qa_summary) if rec.qa_summary else {}
# ---------- 结果 ----------
@app.get("/api/sessions/{sid}/result/preview", response_class=HTMLResponse)
def result_preview(sid: str):
try:
html = service.result_preview(sid)
except ServiceStepError:
raise _error(404, "RESULT_NOT_FOUND", "结果文档不存在")
return HTMLResponse(html)
@app.get("/api/sessions/{sid}/result/download")
def result_download(sid: str):
rec = service.get_session(sid)
if not rec.result_path or not Path(rec.result_path).exists():
raise _error(404, "RESULT_NOT_FOUND", "结果文档不存在")
return FileResponse(rec.result_path, media_type="application/vnd.openxmlformats-officedocument.wordprocessingml.document", filename="output.docx")
@app.get("/api/sessions/{sid}/result/impact-report")
def result_impact(sid: str):
rec = service.get_session(sid)
if not rec.impact_report_path or not Path(rec.impact_report_path).exists():
raise _error(404, "RESULT_NOT_FOUND", "影响调查书不存在")
return FileResponse(rec.impact_report_path, media_type="application/json", filename="impact-report.json")
@app.get("/api/sessions/{sid}/result/qa-report")
def result_qa(sid: str):
rec = service.get_session(sid)
if not rec.qa_report_path or not Path(rec.qa_report_path).exists():
raise _error(404, "RESULT_NOT_FOUND", "QA 报告不存在")
return FileResponse(rec.qa_report_path, media_type="application/json", filename="qa-report.json")
# ---------- 进度流(WebSocket ----------
@app.websocket("/api/sessions/{sid}/ws")
async def session_progress_ws(ws: WebSocket, sid: str):
# 先注册事件循环并订阅,再 accept,缩小「连接已开但服务端尚未订阅」期间
# 的事件丢失窗口(否则客户端连上瞬间触发的进度会被静默丢弃)
hub.register_loop(asyncio.get_running_loop())
q = hub.subscribe(sid)
await ws.accept()
try:
while True:
event = await q.get()
await ws.send_json(event)
except WebSocketDisconnect:
pass
finally:
hub.unsubscribe(sid, q)
# ---------- 聊天 ----------
@app.post("/api/chat/{sid}/messages")
def chat_message(sid: str, body: ChatMessageReq):
try:
result = chat_agent.handle_message(sid, body.content)
except SessionNotFoundError:
raise _error(404, "SESSION_NOT_FOUND", f"会话不存在: {sid}")
return result
@app.get("/api/chat/{sid}/messages")
def chat_history(sid: str):
try:
service.get_session(sid)
except SessionNotFoundError:
raise _error(404, "SESSION_NOT_FOUND", f"会话不存在: {sid}")
return service.store.list_messages(sid)
return app
# 模块级 appuvicorn app.main:app 兼容)
app = create_app()