246 lines
7.2 KiB
Python
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
|