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")