nex_docus/backend/app/api/v1/chat.py

246 lines
7.2 KiB
Python

"""
知识库对话相关 API
"""
from fastapi import APIRouter, Depends, HTTPException, status
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy import select, delete
from typing import List
import json
from app.core.database import get_db
from app.core.deps import get_current_user
from app.models.user import User
from app.models.chat_session import ChatSession, ChatMessage
from app.models.project import Project
from app.models.document_vector import DocumentVector
from app.schemas.response import success_response
from app.services.project_service import get_project_or_404, require_project_read_access
from app.services.rag_service import rag_service
from pydantic import BaseModel
router = APIRouter()
class ChatCreateRequest(BaseModel):
"""创建对话会话请求"""
project_id: int
llm_config_id: int
title: str = "新对话"
class ChatMessageRequest(BaseModel):
"""发送聊天消息请求"""
session_id: int
message: str
class ChatResponse(BaseModel):
"""对话响应"""
session_id: int
user_message: str
assistant_message: str
@router.post("/sessions", response_model=dict)
async def create_chat_session(
req: ChatCreateRequest,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
"""创建新的对话会话"""
project = await get_project_or_404(db, req.project_id)
_, user_role = await require_project_read_access(db, req.project_id, current_user)
session = ChatSession(
user_id=current_user.id,
project_id=req.project_id,
llm_config_id=req.llm_config_id,
title=req.title,
)
db.add(session)
await db.commit()
await db.refresh(session)
return success_response(data={
"session_id": session.id,
"title": session.title,
"created_at": session.created_at.isoformat(),
})
@router.get("/sessions", response_model=dict)
async def list_chat_sessions(
project_id: int,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
"""列出项目的所有对话会话"""
await require_project_read_access(db, project_id, current_user)
stmt = select(ChatSession).where(
ChatSession.project_id == project_id,
ChatSession.user_id == current_user.id,
).order_by(ChatSession.created_at.desc())
result = await db.execute(stmt)
sessions = result.scalars().all()
return success_response(data=[{
"id": s.id,
"title": s.title,
"created_at": s.created_at.isoformat(),
"updated_at": s.updated_at.isoformat(),
} for s in sessions])
@router.get("/sessions/{session_id}/messages", response_model=dict)
async def get_session_messages(
session_id: int,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
"""获取对话会话的所有消息"""
stmt = select(ChatSession).where(ChatSession.id == session_id)
result = await db.execute(stmt)
session = result.scalar_one_or_none()
if not session:
raise HTTPException(status_code=404, detail="对话会话不存在")
if session.user_id != current_user.id:
raise HTTPException(status_code=403, detail="无权访问该对话")
await require_project_read_access(db, session.project_id, current_user)
msg_stmt = select(ChatMessage).where(
ChatMessage.session_id == session_id
).order_by(ChatMessage.created_at.asc())
msg_result = await db.execute(msg_stmt)
messages = msg_result.scalars().all()
return success_response(data=[{
"id": m.id,
"role": m.role,
"content": m.content,
"created_at": m.created_at.isoformat(),
} for m in messages])
@router.post("/send", response_model=dict)
async def send_chat_message(
req: ChatMessageRequest,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
"""发送对话消息并获取回复"""
stmt = select(ChatSession).where(ChatSession.id == req.session_id)
result = await db.execute(stmt)
session = result.scalar_one_or_none()
if not session:
raise HTTPException(status_code=404, detail="对话会话不存在")
if session.user_id != current_user.id:
raise HTTPException(status_code=403, detail="无权访问该对话")
await require_project_read_access(db, session.project_id, current_user)
query_message = ChatMessage(
session_id=req.session_id,
role="user",
content=req.message,
)
db.add(query_message)
await db.flush()
try:
query_vector = await _get_query_vector(req.message)
retrieved_docs = await rag_service.retrieve_documents(
db,
session.project_id,
query_vector,
top_k=5,
)
msg_stmt = select(ChatMessage).where(
ChatMessage.session_id == req.session_id,
).order_by(ChatMessage.created_at.asc())
msg_result = await db.execute(msg_stmt)
prev_messages = msg_result.scalars().all()
conversation_history = [
{"role": m.role, "content": m.content}
for m in prev_messages
if m.id != query_message.id
]
assistant_response = await rag_service.generate_response(
db,
req.message,
query_vector,
session.project_id,
session.llm_config_id,
retrieved_docs,
conversation_history,
)
assistant_message = ChatMessage(
session_id=req.session_id,
role="assistant",
content=assistant_response,
)
db.add(assistant_message)
await db.commit()
return success_response(data={
"session_id": req.session_id,
"user_message": req.message,
"assistant_message": assistant_response,
})
except Exception as e:
await db.rollback()
raise HTTPException(
status_code=500,
detail=f"生成对话回复失败: {str(e)}"
)
@router.delete("/sessions/{session_id}", response_model=dict)
async def delete_chat_session(
session_id: int,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
"""删除对话会话"""
stmt = select(ChatSession).where(ChatSession.id == session_id)
result = await db.execute(stmt)
session = result.scalar_one_or_none()
if not session:
raise HTTPException(status_code=404, detail="对话会话不存在")
if session.user_id != current_user.id:
raise HTTPException(status_code=403, detail="无权删除该对话")
await db.delete(session)
msg_stmt = delete(ChatMessage).where(ChatMessage.session_id == session_id)
await db.execute(msg_stmt)
await db.commit()
return success_response(message="对话会话已删除")
async def _get_query_vector(query: str) -> List[float]:
"""从查询文本生成向量"""
try:
import numpy as np
from sklearn.feature_extraction.text import TfidfVectorizer
vectorizer = TfidfVectorizer(max_features=768)
query_vector = vectorizer.fit_transform([query]).toarray()[0]
return query_vector.tolist()
except Exception:
return [0.0] * 768