ForcePilot/src/routers/chat_router.py
2025-03-18 04:58:20 +08:00

120 lines
4.3 KiB
Python

import json
import asyncio
import traceback
from fastapi import APIRouter, Body
from fastapi.responses import StreamingResponse, Response
from src.core import HistoryManager
from src import executor, config, retriever
from src.models import select_model
from src.utils.logging_config import logger
chat = APIRouter(prefix="/chat")
@chat.get("/")
async def chat_get():
return "Chat Get!"
@chat.post("/")
def chat_post(
query: str = Body(...),
meta: dict = Body(None),
history: list = Body(...),
cur_res_id: str = Body(...)):
model = select_model(config)
meta["server_model_name"] = model.model_name
history_manager = HistoryManager(history)
logger.debug(f"Received query: {query} with meta: {meta}")
def make_chunk(content=None, **kwargs):
return json.dumps({
"response": content,
"model_name": config.model_name,
"meta": meta,
**kwargs
}, ensure_ascii=False).encode('utf-8') + b"\n"
def need_retrieve(meta):
return meta.get("use_web") or meta.get("use_graph") or meta.get("db_name")
def generate_response():
modified_query = query
refs = None
# 处理知识库检索
if meta and need_retrieve(meta):
chunk = make_chunk(status="searching")
yield chunk
try:
modified_query, refs = retriever(modified_query, history_manager.messages, meta)
except Exception as e:
logger.error(f"Retriever error: {e}, {traceback.format_exc()}")
yield make_chunk(message=f"Retriever error: {e}", status="error")
return
yield make_chunk(status="generating")
messages = history_manager.get_history_with_msg(modified_query, max_rounds=meta.get('history_round'))
history_manager.add_user(query) # 注意这里使用原始查询
content = ""
reasoning_content = ""
try:
for delta in model.predict(messages, stream=True):
if not delta.content and hasattr(delta, 'reasoning_content'):
reasoning_content += delta.reasoning_content or ""
chunk = make_chunk(reasoning_content=reasoning_content, status="reasoning")
yield chunk
continue
# 文心一言
if hasattr(delta, 'is_full') and delta.is_full:
content = delta.content
else:
content += delta.content or ""
chunk = make_chunk(content=content, status="loading")
yield chunk
logger.debug(f"Final response: {content}")
logger.debug(f"Final reasoning response: {reasoning_content}")
yield make_chunk(content=content,
status="finished",
history=history_manager.update_ai(content),
refs=refs)
except Exception as e:
logger.error(f"Model error: {e}, {traceback.format_exc()}")
yield make_chunk(message=f"Model error: {e}", status="error")
return
return StreamingResponse(generate_response(), media_type='application/json')
@chat.post("/call")
async def call(query: str = Body(...), meta: dict = Body(None)):
model = select_model(config, model_provider=meta.get("model_provider"), model_name=meta.get("model_name"))
async def predict_async(query):
loop = asyncio.get_event_loop()
return await loop.run_in_executor(executor, model.predict, query)
response = await predict_async(query)
logger.debug({"query": query, "response": response.content})
return {"response": response.content}
@chat.post("/call_lite")
async def call(query: str = Body(...), meta: dict = Body(None)):
meta = meta or {}
async def predict_async(query):
loop = asyncio.get_event_loop()
model_provider = meta.get("model_provider", config.model_provider_lite)
model_name = meta.get("model_name", config.model_name_lite)
model = select_model(config, model_provider=model_provider, model_name=model_name)
return await loop.run_in_executor(executor, model.predict, query)
response = await predict_async(query)
logger.debug({"query": query, "response": response.content})
return {"response": response.content}