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 }