- 将 /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% 达标
355 lines
13 KiB
Python
355 lines
13 KiB
Python
"""Web 服务化:FastAPI 端点(S4)。
|
||
|
||
实现 api-design.md §2 的核心端点(v1 同步执行 + 内嵌零构建前端)。
|
||
- 会话:POST/GET /api/sessions、GET /api/sessions/{id}、DELETE
|
||
- 文件:POST /api/sessions/{id}/files(multipart)
|
||
- 解析:POST start-parse、GET parse-result、POST confirm-parse
|
||
- 影响:POST start-impact、GET impact-result、POST confirm-impact
|
||
- 生成/QA:POST 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
|
||
|
||
|
||
# 模块级 app(uvicorn app.main:app 兼容)
|
||
app = create_app()
|