projects/repobot/v4/server.py
98 行 · 3.2 KB程式碼和執行結果保留原樣(簡體中文),與實際執行時完全一致。
"""RepoBot v4 的网页服务。
uvicorn server:app --host 127.0.0.1 --port 8000
然后打开 http://127.0.0.1:8000
接口:POST /api/chat,请求体 {"message": "问题", "history": [之前的 user/assistant 消息]},
返回 SSE 事件流,每个事件是一行 "data: {json}"。
服务器不保存对话,历史由前端带上来。
"""
import json
from contextlib import asynccontextmanager
from pathlib import Path
from fastapi import FastAPI, HTTPException
from fastapi.responses import FileResponse, StreamingResponse
from pydantic import BaseModel, Field
from starlette.concurrency import run_in_threadpool
import agent
import guard
import llm
import tools
from retrieval import Retriever
from tracing import Tracer
HERE = Path(__file__).parent
tracer = Tracer(HERE / "logs" / "traces.jsonl")
MAX_HISTORY = 10
@asynccontextmanager
async def lifespan(app):
# 启动时加载一次模型和索引,所有请求共用
(HERE / "logs").mkdir(exist_ok=True)
tools.ensure_source()
tools.retriever = Retriever(tools.DOCS_DIR, HERE / ".cache")
yield
app = FastAPI(lifespan=lifespan)
class Message(BaseModel):
role: str = Field(pattern="^(user|assistant)$") # 不许前端塞进 system 或 tool 消息
content: str = Field(max_length=8000)
class ChatRequest(BaseModel):
message: str = Field(min_length=1, max_length=2000)
history: list[Message] = []
def sse(event):
return f"data: {json.dumps(event, ensure_ascii=False)}\n\n"
@app.post("/api/chat")
async def chat(req: ChatRequest):
history = [m.model_dump() for m in req.history[-MAX_HISTORY:]]
label, usage = await run_in_threadpool(guard.classify, req.message)
def events():
with tracer.span("task", "chat", question=req.message[:200], guard=label) as task:
if label != "httpx":
task["cost"] = round(llm.cost_usd(usage), 6)
yield sse({"type": "token", "text": guard.REPLIES[label]})
yield sse({"type": "done", "steps": 0, "tool_calls": 0, "cost": task["cost"]})
return
redactor = guard.LineRedactor()
for event in agent.run_stream(req.message, history, tracer):
if event["type"] == "token":
text = redactor.feed(event["text"])
if text:
yield sse({"type": "token", "text": text})
continue
if event["type"] == "done":
rest = redactor.flush()
if rest:
yield sse({"type": "token", "text": rest})
event["cost"] = round(event["cost"] + llm.cost_usd(usage), 6)
task.update(steps=event["steps"], cost=event["cost"])
yield sse(event)
# events 是普通的生成器,StreamingResponse 会把它放到线程池里执行,不会卡住服务器
return StreamingResponse(events(), media_type="text/event-stream")
@app.get("/healthz")
async def healthz():
if tools.retriever is None:
raise HTTPException(503, "还在加载")
return {"ok": True}
@app.get("/")
async def index():
return FileResponse(HERE / "static" / "index.html")