diff --git a/src/main.py b/src/main.py index 18a97011..d831e3d0 100644 --- a/src/main.py +++ b/src/main.py @@ -24,5 +24,5 @@ logger = setup_logger("server:main") if __name__ == "__main__": - uvicorn.run(app, host="0.0.0.0", port=5000) + uvicorn.run(app, host="0.0.0.0", port=5000, threads=10, workers=10) diff --git a/src/routers/base_router.py b/src/routers/base_router.py index dda125a5..31ac506c 100644 --- a/src/routers/base_router.py +++ b/src/routers/base_router.py @@ -16,7 +16,7 @@ async def route_index(): return {"message": "You Got It!"} @base.get("/config") -async def get_config(): +def get_config(): return startup.config @base.post("/config") @@ -32,7 +32,7 @@ async def restart(): return {"message": "Restarted!"} @base.get("/log") -async def get_log(): +def get_log(): from src.utils.logging_config import LOG_FILE from collections import deque diff --git a/src/routers/chat_router.py b/src/routers/chat_router.py index d6dcc0d5..f7a9271c 100644 --- a/src/routers/chat_router.py +++ b/src/routers/chat_router.py @@ -1,31 +1,41 @@ -from fastapi import APIRouter, HTTPException, Request -from fastapi.responses import StreamingResponse +import time +import uuid +from fastapi import APIRouter, HTTPException, Request, Body +from fastapi.responses import StreamingResponse, Response +from concurrent.futures import ThreadPoolExecutor import json +import asyncio from src.core import HistoryManager from src.core.startup import startup from src.utils.logging_config import setup_logger chat = APIRouter(prefix="/chat") logger = setup_logger("server-chat") +# 创建线程池 +executor = ThreadPoolExecutor() + +refs_pool = {} @chat.get("/") async def chat_get(): return "Chat Get!" @chat.post("/") -async def chat_post(request: Request): - request_data = await request.json() - query = request_data['query'] - meta = request_data.get('meta') - history_manager = HistoryManager(request_data['history']) +def chat_post( + query: str = Body(...), + meta: dict = Body(None), + history: list = Body(...), + cur_res_id: str = Body(...)): + history_manager = HistoryManager(history) new_query, refs = startup.retriever(query, history_manager.messages, meta) + refs_pool[cur_res_id] = refs messages = history_manager.get_history_with_msg(new_query, max_rounds=meta.get('history_round')) history_manager.add_user(query) logger.debug(f"Web history: {history_manager.messages}") - async def generate_response(): + def generate_response(): content = "" for delta in startup.model.predict(messages, stream=True): if not delta.content: @@ -36,20 +46,29 @@ async def chat_post(request: Request): else: content += delta.content - response_chunk = json.dumps({ - "history": history_manager.update_ai(content), + logger.debug(f"Response: {content}") + + _chunk = json.dumps({ "response": content, - "refs": refs # TODO: 优化 refs,不需要每次都返回 - }, ensure_ascii=False).encode('utf8') + b'\n' - yield response_chunk + "history": history_manager.update_ai(content), + }, ensure_ascii=False).encode('utf-8') + b"\n" + yield _chunk return StreamingResponse(generate_response(), media_type='application/json') @chat.post("/call") -async def call(request: Request): - request_data = await request.json() - query = request_data['query'] - response = startup.model.predict(query) +async def call(query: str = Body(...), meta: dict = Body(None)): + async def predict_async(query): + loop = asyncio.get_event_loop() + return await loop.run_in_executor(executor, startup.model.predict, query) + + response = await predict_async(query) logger.debug({"query": query, "response": response.content}) - return {"response": response.content} \ No newline at end of file + return {"response": response.content} + +@chat.get("/refs") +def get_refs(cur_res_id: str): + global refs_pool + refs = refs_pool.pop(cur_res_id, None) + return {"refs": refs} \ No newline at end of file diff --git a/src/routers/data_router.py b/src/routers/data_router.py index f5073f98..2a19805b 100644 --- a/src/routers/data_router.py +++ b/src/routers/data_router.py @@ -1,7 +1,6 @@ import os from typing import List, Optional from fastapi import APIRouter, File, UploadFile, HTTPException, Depends, Body -from pydantic import BaseModel from src.utils import setup_logger, hashstr from src.core.startup import startup diff --git a/web/src/components/ChatComponent.vue b/web/src/components/ChatComponent.vue index 504834ea..40d84b8c 100644 --- a/web/src/components/ChatComponent.vue +++ b/web/src/components/ChatComponent.vue @@ -16,10 +16,9 @@