""" 知识库对话相关 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