ForcePilot/server/routers/chat_router.py

405 lines
14 KiB
Python
Raw Normal View History

import asyncio
import json
2025-03-17 19:58:00 +08:00
import traceback
2025-03-25 05:40:07 +08:00
import uuid
import yaml
from pathlib import Path
from fastapi import APIRouter, Body, Depends, HTTPException
2025-03-25 05:40:07 +08:00
from fastapi.responses import StreamingResponse
from langchain_core.messages import AIMessageChunk, HumanMessage
from pydantic import BaseModel
from sqlalchemy.orm import Session
2025-03-25 05:40:07 +08:00
from server.models.thread_model import Thread
from server.models.user_model import User
from server.routers.auth_router import get_admin_user
from server.utils.auth_middleware import get_db, get_required_user
from src import config, executor
2025-03-25 05:40:07 +08:00
from src.agents import agent_manager
from src.agents.common.tools import gen_tool_info, get_buildin_tools
2025-03-16 01:11:13 +08:00
from src.models import select_model
2025-02-27 19:35:25 +08:00
from src.utils.logging_config import logger
2024-10-02 20:11:28 +08:00
2025-07-22 17:29:38 +08:00
chat = APIRouter(prefix="/chat", tags=["chat"])
# =============================================================================
# > === 智能体管理分组 ===
# =============================================================================
2025-03-06 23:55:20 +08:00
2025-05-02 23:56:59 +08:00
@chat.get("/default_agent")
async def get_default_agent(current_user: User = Depends(get_required_user)):
"""获取默认智能体ID需要登录"""
try:
default_agent_id = config.default_agent_id
# 如果没有设置默认智能体,尝试获取第一个可用的智能体
if not default_agent_id:
agents = await agent_manager.get_agents_info()
2025-05-02 23:56:59 +08:00
if agents:
default_agent_id = agents[0].get("id", "")
2025-05-02 23:56:59 +08:00
return {"default_agent_id": default_agent_id}
except Exception as e:
logger.error(f"获取默认智能体出错: {e}")
raise HTTPException(status_code=500, detail=f"获取默认智能体出错: {str(e)}")
2025-05-02 23:56:59 +08:00
@chat.post("/set_default_agent")
async def set_default_agent(request_data: dict = Body(...), current_user=Depends(get_admin_user)):
2025-05-02 23:56:59 +08:00
"""设置默认智能体ID (仅管理员)"""
try:
agent_id = request_data.get("agent_id")
if not agent_id:
raise HTTPException(status_code=422, detail="缺少必需的 agent_id 字段")
2025-05-02 23:56:59 +08:00
# 验证智能体是否存在
agents = await agent_manager.get_agents_info()
agent_ids = [agent.get("id", "") for agent in agents]
2025-05-02 23:56:59 +08:00
if agent_id not in agent_ids:
raise HTTPException(status_code=404, detail=f"智能体 {agent_id} 不存在")
# 设置默认智能体ID
config.default_agent_id = agent_id
# 保存配置
config.save()
return {"success": True, "default_agent_id": agent_id}
except HTTPException as he:
raise he
except Exception as e:
logger.error(f"设置默认智能体出错: {e}")
raise HTTPException(status_code=500, detail=f"设置默认智能体出错: {str(e)}")
2025-07-22 17:29:38 +08:00
# =============================================================================
# > === 对话分组 ===
# =============================================================================
2024-10-02 20:11:28 +08:00
@chat.post("/call")
2025-05-02 23:56:59 +08:00
async def call(query: str = Body(...), meta: dict = Body(None), current_user: User = Depends(get_required_user)):
"""调用模型进行简单问答(需要登录)"""
2025-04-10 11:41:01 +08:00
meta = meta or {}
2025-03-29 17:33:09 +08:00
model = select_model(model_provider=meta.get("model_provider"), model_name=meta.get("model_name"))
async def predict_async(query):
loop = asyncio.get_event_loop()
2025-03-16 01:11:13 +08:00
return await loop.run_in_executor(executor, model.predict, query)
response = await predict_async(query)
2024-10-02 20:11:28 +08:00
logger.debug({"query": query, "response": response.content})
return {"response": response.content}
2025-03-25 05:40:07 +08:00
@chat.get("/agent")
2025-05-02 23:56:59 +08:00
async def get_agent(current_user: User = Depends(get_required_user)):
"""获取所有可用智能体(需要登录)"""
agents = await agent_manager.get_agents_info()
# logger.debug(f"agents: {agents}")
metadata = {}
if Path("src/static/agents_meta.yaml").exists():
with open("src/static/agents_meta.yaml") as f:
metadata = yaml.safe_load(f)
return {"agents": agents, "metadata": metadata}
2025-03-25 05:40:07 +08:00
@chat.post("/agent/{agent_id}")
async def chat_agent(
agent_id: str,
query: str = Body(...),
config: dict = Body({}),
meta: dict = Body({}),
current_user: User = Depends(get_required_user),
):
2025-05-02 23:56:59 +08:00
"""使用特定智能体进行对话(需要登录)"""
2025-04-02 13:00:25 +08:00
logger.info(f"agent_id: {agent_id}, query: {query}, config: {config}, meta: {meta}")
meta.update(
{
"query": query,
"agent_id": agent_id,
"server_model_name": config.get("model", agent_id),
"thread_id": config.get("thread_id"),
"user_id": current_user.id,
}
)
2025-04-02 00:00:04 +08:00
# 将meta和thread_id整合到config中
def make_chunk(content=None, **kwargs):
return (
json.dumps(
{"request_id": meta.get("request_id"), "response": content, **kwargs}, ensure_ascii=False
).encode("utf-8")
+ b"\n"
)
2025-05-16 23:46:25 +08:00
async def stream_messages():
# 代表服务端已经收到了请求
yield make_chunk(status="init", meta=meta, msg=HumanMessage(content=query).model_dump())
try:
agent = agent_manager.get_agent(agent_id)
except Exception as e:
logger.error(f"Error getting agent {agent_id}: {e}, {traceback.format_exc()}")
yield make_chunk(message=f"Error getting agent {agent_id}: {e}", status="error")
return
messages = [{"role": "user", "content": query}]
# 构造运行时配置如果没有thread_id则生成一个
user_id = str(current_user.id)
thread_id = config.get("thread_id")
input_context = {"user_id": user_id, "thread_id": thread_id}
2025-04-05 02:23:32 +08:00
try:
async for msg, metadata in agent.stream_messages(messages, input_context=input_context):
# logger.debug(f"msg: {msg.model_dump()}, metadata: {metadata}")
if isinstance(msg, AIMessageChunk):
yield make_chunk(content=msg.content, msg=msg.model_dump(), metadata=metadata, status="loading")
2025-04-05 02:23:32 +08:00
else:
yield make_chunk(msg=msg.model_dump(), metadata=metadata, status="loading")
2025-04-05 02:23:32 +08:00
yield make_chunk(status="finished", meta=meta)
2025-04-05 02:23:32 +08:00
except Exception as e:
logger.error(f"Error streaming messages: {e}, {traceback.format_exc()}")
2025-04-05 02:23:32 +08:00
yield make_chunk(message=f"Error streaming messages: {e}", status="error")
2025-03-25 05:40:07 +08:00
return StreamingResponse(stream_messages(), media_type="application/json")
2025-04-01 00:39:54 +08:00
2025-07-22 17:29:38 +08:00
# =============================================================================
# > === 模型管理分组 ===
# =============================================================================
2025-04-01 00:39:54 +08:00
@chat.get("/models")
2025-05-02 23:56:59 +08:00
async def get_chat_models(model_provider: str, current_user: User = Depends(get_admin_user)):
"""获取指定模型提供商的模型列表(需要登录)"""
2025-04-01 00:39:54 +08:00
model = select_model(model_provider=model_provider)
return {"models": model.get_models()}
@chat.post("/models/update")
async def update_chat_models(model_provider: str, model_names: list[str], current_user=Depends(get_admin_user)):
2025-05-02 23:56:59 +08:00
"""更新指定模型提供商的模型列表 (仅管理员)"""
config.model_names[model_provider]["models"] = model_names
config._save_models_to_file()
return {"models": config.model_names[model_provider]["models"]}
@chat.get("/tools")
async def get_tools(agent_id: str, current_user: User = Depends(get_required_user)):
2025-05-02 23:56:59 +08:00
"""获取所有可用工具(需要登录)"""
# 获取Agent实例和配置类
if not (agent := agent_manager.get_agent(agent_id)):
raise HTTPException(status_code=404, detail=f"智能体 {agent_id} 不存在")
if hasattr(agent, "get_tools"):
tools = agent.get_tools()
else:
tools = get_buildin_tools()
tools_info = gen_tool_info(tools)
return {"tools": {tool["id"]: tool for tool in tools_info}}
@chat.post("/agent/{agent_id}/config")
async def save_agent_config(agent_id: str, config: dict = Body(...), current_user: User = Depends(get_required_user)):
"""保存智能体配置到YAML文件需要登录"""
try:
# 获取Agent实例和配置类
if not (agent := agent_manager.get_agent(agent_id)):
raise HTTPException(status_code=404, detail=f"智能体 {agent_id} 不存在")
# 使用配置类的save_to_file方法保存配置
result = agent.context_schema.save_to_file(config, agent.module_name)
if result:
return {"success": True, "message": f"智能体 {agent.name} 配置已保存"}
else:
2025-05-24 11:29:45 +08:00
raise HTTPException(status_code=500, detail="保存智能体配置失败")
except Exception as e:
logger.error(f"保存智能体配置出错: {e}, {traceback.format_exc()}")
raise HTTPException(status_code=500, detail=f"保存智能体配置出错: {str(e)}")
@chat.get("/agent/{agent_id}/history")
async def get_agent_history(agent_id: str, thread_id: str, current_user: User = Depends(get_required_user)):
2025-05-15 23:45:14 +08:00
"""获取智能体历史消息(需要登录)"""
try:
# 获取Agent实例和配置类
if not (agent := agent_manager.get_agent(agent_id)):
raise HTTPException(status_code=404, detail=f"智能体 {agent_id} 不存在")
# 获取历史消息
history = await agent.get_history(user_id=str(current_user.id), thread_id=thread_id)
return {"history": history}
except Exception as e:
logger.error(f"获取智能体历史消息出错: {e}, {traceback.format_exc()}")
raise HTTPException(status_code=500, detail=f"获取智能体历史消息出错: {str(e)}")
@chat.get("/agent/{agent_id}/config")
async def get_agent_config(agent_id: str, current_user: User = Depends(get_required_user)):
2025-05-15 23:45:14 +08:00
"""从YAML文件加载智能体配置需要登录"""
try:
# 检查智能体是否存在
if not (agent := agent_manager.get_agent(agent_id)):
raise HTTPException(status_code=404, detail=f"智能体 {agent_id} 不存在")
config = await agent.get_config()
logger.debug(f"config: {config}, ContextClass: {agent.context_schema=}")
return {"success": True, "config": config}
except Exception as e:
logger.error(f"加载智能体配置出错: {e}, {traceback.format_exc()}")
raise HTTPException(status_code=500, detail=f"加载智能体配置出错: {str(e)}")
# ==================== 线程管理 API ====================
class ThreadCreate(BaseModel):
2025-05-24 11:29:45 +08:00
title: str | None = None
agent_id: str
2025-05-24 11:29:45 +08:00
description: str | None = None
metadata: dict | None = None
class ThreadResponse(BaseModel):
id: str
user_id: str
agent_id: str
2025-05-24 11:29:45 +08:00
title: str | None = None
description: str | None = None
create_at: str
update_at: str
2025-07-22 17:29:38 +08:00
# =============================================================================
# > === 会话管理分组 ===
# =============================================================================
@chat.post("/thread", response_model=ThreadResponse)
async def create_thread(
thread: ThreadCreate, db: Session = Depends(get_db), current_user: User = Depends(get_required_user)
):
"""创建新对话线程"""
thread_id = str(uuid.uuid4())
logger.debug(f"thread.agent_id: {thread.agent_id}")
new_thread = Thread(
id=thread_id,
user_id=str(current_user.id),
agent_id=thread.agent_id,
title=thread.title or "新的对话",
description=thread.description,
)
db.add(new_thread)
db.commit()
db.refresh(new_thread)
return {
"id": new_thread.id,
"user_id": new_thread.user_id,
"agent_id": new_thread.agent_id,
"title": new_thread.title,
"description": new_thread.description,
"create_at": new_thread.create_at.isoformat(),
"update_at": new_thread.update_at.isoformat(),
}
2025-05-24 11:29:45 +08:00
@chat.get("/threads", response_model=list[ThreadResponse])
async def list_threads(agent_id: str, db: Session = Depends(get_db), current_user: User = Depends(get_required_user)):
"""获取用户的所有对话线程"""
assert agent_id, "agent_id 不能为空"
query = db.query(Thread).filter(
Thread.user_id == str(current_user.id),
Thread.status == 1,
Thread.agent_id == agent_id,
)
logger.debug(f"agent_id: {agent_id}")
threads = query.order_by(Thread.update_at.desc()).all()
return [
{
"id": thread.id,
"user_id": thread.user_id,
"agent_id": thread.agent_id,
"title": thread.title,
"description": thread.description,
"create_at": thread.create_at.isoformat(),
"update_at": thread.update_at.isoformat(),
}
for thread in threads
]
@chat.delete("/thread/{thread_id}")
async def delete_thread(thread_id: str, db: Session = Depends(get_db), current_user: User = Depends(get_required_user)):
"""删除对话线程"""
thread = db.query(Thread).filter(Thread.id == thread_id, Thread.user_id == str(current_user.id)).first()
if not thread:
raise HTTPException(status_code=404, detail="对话线程不存在")
# 软删除
thread.status = 0
db.commit()
return {"message": "删除成功"}
class ThreadUpdate(BaseModel):
2025-05-24 11:29:45 +08:00
title: str | None = None
description: str | None = None
@chat.put("/thread/{thread_id}", response_model=ThreadResponse)
async def update_thread(
thread_id: str,
thread_update: ThreadUpdate,
db: Session = Depends(get_db),
current_user: User = Depends(get_required_user),
):
"""更新对话线程信息"""
thread = (
db.query(Thread)
.filter(Thread.id == thread_id, Thread.user_id == str(current_user.id), Thread.status == 1)
.first()
)
if not thread:
raise HTTPException(status_code=404, detail="对话线程不存在")
if thread_update.title is not None:
thread.title = thread_update.title
if thread_update.description is not None:
thread.description = thread_update.description
db.commit()
db.refresh(thread)
return {
"id": thread.id,
"user_id": thread.user_id,
"agent_id": thread.agent_id,
"title": thread.title,
"description": thread.description,
"create_at": thread.create_at.isoformat(),
"update_at": thread.update_at.isoformat(),
}