278 lines
7.5 KiB
Go
278 lines
7.5 KiB
Go
package api
|
|||
|
|
|
||
|
|
import (
|
||
|
|
"net/http"
|
||
|
|
"time"
|
||
|
|
|
||
|
|
"d:\项目\dobaochet\backend\config"
|
||
|
|
"d:\项目\dobaochet\backend\db"
|
||
|
|
"d:\项目\dobaochet\backend\models"
|
||
|
|
"d:\项目\dobaochet\backend\services"
|
||
|
|
|
||
|
|
"github.com/gin-gonic/gin"
|
||
|
|
)
|
||
|
|
|
||
|
|
// ChatAPI 聊天API结构体
|
||
|
|
type ChatAPI struct {
|
||
|
|
cfg *config.Config
|
||
|
|
aiService *services.AIService
|
||
|
|
}
|
||
|
|
|
||
|
|
// NewChatAPI 创建聊天API实例
|
||
|
|
func NewChatAPI(cfg *config.Config, aiService *services.AIService) *ChatAPI {
|
||
|
|
return &ChatAPI{
|
||
|
|
cfg: cfg,
|
||
|
|
aiService: aiService,
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
// Chat 处理聊天请求
|
||
|
|
func (api *ChatAPI) Chat(c *gin.Context) {
|
||
|
|
// 从上下文获取用户信息
|
||
|
|
user, exists := c.Get("user")
|
||
|
|
if !exists {
|
||
|
|
c.JSON(http.StatusUnauthorized, gin.H{"error": "未认证"})
|
||
|
|
return
|
||
|
|
}
|
||
|
|
|
||
|
|
currentUser := user.(models.User)
|
||
|
|
|
||
|
|
// 解析请求
|
||
|
|
var req models.ChatRequest
|
||
|
|
if err := c.ShouldBindJSON(&req); err != nil {
|
||
|
|
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||
|
|
return
|
||
|
|
}
|
||
|
|
|
||
|
|
// 估算用户输入的令牌数量
|
||
|
|
inputTokens := api.aiService.CountTokens(req.Message)
|
||
|
|
|
||
|
|
// 检查单次请求令牌限制
|
||
|
|
if inputTokens > currentUser.Quota.TokenLimit {
|
||
|
|
c.JSON(http.StatusBadRequest, gin.H{"error": "单次请求令牌数超过限制"})
|
||
|
|
return
|
||
|
|
}
|
||
|
|
|
||
|
|
// 检查配额是否足够
|
||
|
|
if currentUser.Quota.UsedTokens+inputTokens > currentUser.Quota.TotalTokens {
|
||
|
|
c.JSON(http.StatusForbidden, gin.H{"error": "令牌配额不足"})
|
||
|
|
return
|
||
|
|
}
|
||
|
|
|
||
|
|
// 开始数据库事务
|
||
|
|
tx := db.GetDB().Begin()
|
||
|
|
|
||
|
|
// 获取或创建对话
|
||
|
|
var conversation models.Conversation
|
||
|
|
if req.ConversationID > 0 {
|
||
|
|
// 检查对话是否属于当前用户
|
||
|
|
if result := tx.Where("id = ? AND user_id = ?", req.ConversationID, currentUser.ID).First(&conversation); result.Error != nil {
|
||
|
|
tx.Rollback()
|
||
|
|
c.JSON(http.StatusNotFound, gin.H{"error": "对话不存在"})
|
||
|
|
return
|
||
|
|
}
|
||
|
|
} else {
|
||
|
|
// 创建新对话
|
||
|
|
conversation = models.Conversation{
|
||
|
|
UserID: currentUser.ID,
|
||
|
|
Title: req.Message[:min(50, len(req.Message))],
|
||
|
|
TotalTokens: 0,
|
||
|
|
}
|
||
|
|
if result := tx.Create(&conversation); result.Error != nil {
|
||
|
|
tx.Rollback()
|
||
|
|
c.JSON(http.StatusInternalServerError, gin.H{"error": "创建对话失败"})
|
||
|
|
return
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
// 保存用户消息
|
||
|
|
userMessage := models.Message{
|
||
|
|
ConversationID: conversation.ID,
|
||
|
|
Role: "user",
|
||
|
|
Content: req.Message,
|
||
|
|
Tokens: inputTokens,
|
||
|
|
Model: api.cfg.AI.ModelName,
|
||
|
|
}
|
||
|
|
if result := tx.Create(&userMessage); result.Error != nil {
|
||
|
|
tx.Rollback()
|
||
|
|
c.JSON(http.StatusInternalServerError, gin.H{"error": "保存消息失败"})
|
||
|
|
return
|
||
|
|
}
|
||
|
|
|
||
|
|
// 获取对话历史(最近10条消息)
|
||
|
|
var historyMessages []models.Message
|
||
|
|
tx.Where("conversation_id = ?", conversation.ID).Order("created_at desc").Limit(10).Find(&historyMessages)
|
||
|
|
|
||
|
|
// 转换为AI服务需要的格式
|
||
|
|
var aiMessages []services.Message
|
||
|
|
for i := len(historyMessages) - 1; i >= 0; i-- {
|
||
|
|
aiMessages = append(aiMessages, services.Message{
|
||
|
|
Role: historyMessages[i].Role,
|
||
|
|
Content: historyMessages[i].Content,
|
||
|
|
})
|
||
|
|
}
|
||
|
|
|
||
|
|
// 调用AI服务
|
||
|
|
aiResp, err := api.aiService.Chat(aiMessages)
|
||
|
|
if err != nil {
|
||
|
|
tx.Rollback()
|
||
|
|
c.JSON(http.StatusInternalServerError, gin.H{"error": "AI服务调用失败: " + err.Error()})
|
||
|
|
return
|
||
|
|
}
|
||
|
|
|
||
|
|
// 检查AI响应
|
||
|
|
if len(aiResp.Choices) == 0 {
|
||
|
|
tx.Rollback()
|
||
|
|
c.JSON(http.StatusInternalServerError, gin.H{"error": "AI服务返回空响应"})
|
||
|
|
return
|
||
|
|
}
|
||
|
|
|
||
|
|
// 保存AI回复
|
||
|
|
aiMessage := models.Message{
|
||
|
|
ConversationID: conversation.ID,
|
||
|
|
Role: "assistant",
|
||
|
|
Content: aiResp.Choices[0].Message.Content,
|
||
|
|
Tokens: aiResp.Usage.CompletionTokens,
|
||
|
|
Model: api.cfg.AI.ModelName,
|
||
|
|
}
|
||
|
|
if result := tx.Create(&aiMessage); result.Error != nil {
|
||
|
|
tx.Rollback()
|
||
|
|
c.JSON(http.StatusInternalServerError, gin.H{"error": "保存AI回复失败"})
|
||
|
|
return
|
||
|
|
}
|
||
|
|
|
||
|
|
// 更新对话总令牌数
|
||
|
|
conversation.TotalTokens += aiResp.Usage.TotalTokens
|
||
|
|
if result := tx.Save(&conversation); result.Error != nil {
|
||
|
|
tx.Rollback()
|
||
|
|
c.JSON(http.StatusInternalServerError, gin.H{"error": "更新对话令牌数失败"})
|
||
|
|
return
|
||
|
|
}
|
||
|
|
|
||
|
|
// 更新用户配额
|
||
|
|
currentUser.Quota.UsedTokens += aiResp.Usage.TotalTokens
|
||
|
|
if result := tx.Save(¤tUser.Quota); result.Error != nil {
|
||
|
|
tx.Rollback()
|
||
|
|
c.JSON(http.StatusInternalServerError, gin.H{"error": "更新配额失败"})
|
||
|
|
return
|
||
|
|
}
|
||
|
|
|
||
|
|
// 提交事务
|
||
|
|
if err := tx.Commit().Error; err != nil {
|
||
|
|
c.JSON(http.StatusInternalServerError, gin.H{"error": "保存数据失败"})
|
||
|
|
return
|
||
|
|
}
|
||
|
|
|
||
|
|
// 返回响应
|
||
|
|
c.JSON(http.StatusOK, models.ChatResponse{
|
||
|
|
Response: aiMessage.Content,
|
||
|
|
ConversationID: conversation.ID,
|
||
|
|
TokensUsed: aiResp.Usage.TotalTokens,
|
||
|
|
})
|
||
|
|
}
|
||
|
|
|
||
|
|
// GetConversations 获取用户对话列表
|
||
|
|
func (api *ChatAPI) GetConversations(c *gin.Context) {
|
||
|
|
// 从上下文获取用户信息
|
||
|
|
userID, exists := c.Get("user_id")
|
||
|
|
if !exists {
|
||
|
|
c.JSON(http.StatusUnauthorized, gin.H{"error": "未认证"})
|
||
|
|
return
|
||
|
|
}
|
||
|
|
|
||
|
|
// 获取对话列表
|
||
|
|
var conversations []models.Conversation
|
||
|
|
db.GetDB().Where("user_id = ?", userID).Order("created_at desc").Find(&conversations)
|
||
|
|
|
||
|
|
c.JSON(http.StatusOK, conversations)
|
||
|
|
}
|
||
|
|
|
||
|
|
// GetConversation 获取对话详情
|
||
|
|
func (api *ChatAPI) GetConversation(c *gin.Context) {
|
||
|
|
// 从上下文获取用户信息
|
||
|
|
userID, exists := c.Get("user_id")
|
||
|
|
if !exists {
|
||
|
|
c.JSON(http.StatusUnauthorized, gin.H{"error": "未认证"})
|
||
|
|
return
|
||
|
|
}
|
||
|
|
|
||
|
|
// 获取对话ID
|
||
|
|
id := c.Param("id")
|
||
|
|
|
||
|
|
// 检查对话是否属于当前用户
|
||
|
|
var conversation models.Conversation
|
||
|
|
if result := db.GetDB().Preload("Messages").Where("id = ? AND user_id = ?", id, userID).First(&conversation); result.Error != nil {
|
||
|
|
c.JSON(http.StatusNotFound, gin.H{"error": "对话不存在"})
|
||
|
|
return
|
||
|
|
}
|
||
|
|
|
||
|
|
c.JSON(http.StatusOK, conversation)
|
||
|
|
}
|
||
|
|
|
||
|
|
// DeleteConversation 删除对话
|
||
|
|
func (api *ChatAPI) DeleteConversation(c *gin.Context) {
|
||
|
|
// 从上下文获取用户信息
|
||
|
|
userID, exists := c.Get("user_id")
|
||
|
|
if !exists {
|
||
|
|
c.JSON(http.StatusUnauthorized, gin.H{"error": "未认证"})
|
||
|
|
return
|
||
|
|
}
|
||
|
|
|
||
|
|
// 获取对话ID
|
||
|
|
id := c.Param("id")
|
||
|
|
|
||
|
|
// 检查对话是否属于当前用户
|
||
|
|
var conversation models.Conversation
|
||
|
|
if result := db.GetDB().Where("id = ? AND user_id = ?", id, userID).First(&conversation); result.Error != nil {
|
||
|
|
c.JSON(http.StatusNotFound, gin.H{"error": "对话不存在"})
|
||
|
|
return
|
||
|
|
}
|
||
|
|
|
||
|
|
// 删除对话(级联删除消息)
|
||
|
|
if result := db.GetDB().Delete(&conversation); result.Error != nil {
|
||
|
|
c.JSON(http.StatusInternalServerError, gin.H{"error": "删除对话失败"})
|
||
|
|
return
|
||
|
|
}
|
||
|
|
|
||
|
|
c.JSON(http.StatusOK, gin.H{"message": "对话删除成功"})
|
||
|
|
}
|
||
|
|
|
||
|
|
// ClearConversations 清空用户对话
|
||
|
|
func (api *ChatAPI) ClearConversations(c *gin.Context) {
|
||
|
|
// 从上下文获取用户信息
|
||
|
|
userID, exists := c.Get("user_id")
|
||
|
|
if !exists {
|
||
|
|
c.JSON(http.StatusUnauthorized, gin.H{"error": "未认证"})
|
||
|
|
return
|
||
|
|
}
|
||
|
|
|
||
|
|
// 删除用户所有对话
|
||
|
|
if result := db.GetDB().Where("user_id = ?", userID).Delete(&models.Conversation{}); result.Error != nil {
|
||
|
|
c.JSON(http.StatusInternalServerError, gin.H{"error": "清空对话失败"})
|
||
|
|
return
|
||
|
|
}
|
||
|
|
|
||
|
|
c.JSON(http.StatusOK, gin.H{"message": "对话清空成功"})
|
||
|
|
}
|
||
|
|
|
||
|
|
// RegisterRoutes 注册聊天路由
|
||
|
|
func (api *ChatAPI) RegisterRoutes(router *gin.RouterGroup) {
|
||
|
|
chatGroup := router.Group("/chat")
|
||
|
|
chatGroup.Use(AuthMiddleware(api.cfg))
|
||
|
|
{
|
||
|
|
chatGroup.POST("/", api.Chat)
|
||
|
|
chatGroup.GET("/conversations", api.GetConversations)
|
||
|
|
chatGroup.GET("/conversations/:id", api.GetConversation)
|
||
|
|
chatGroup.DELETE("/conversations/:id", api.DeleteConversation)
|
||
|
|
chatGroup.DELETE("/conversations", api.ClearConversations)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
// min 返回两个整数中的最小值
|
||
|
|
func min(a, b int) int {
|
||
|
|
if a < b {
|
||
|
|
return a
|
||
|
|
}
|
||
|
|
return b
|
||
|
|
}
|