From 2c8afc1275fdcb7cd92ae1e839d0d1ac3d68c6e7 Mon Sep 17 00:00:00 2001 From: sun <3371392206@qq.com> Date: Fri, 2 Jan 2026 01:19:24 +0800 Subject: [PATCH] =?UTF-8?q?=E4=B8=8A=E4=BC=A0=E5=88=9D=E5=A7=8B=E7=89=88?= =?UTF-8?q?=EF=BC=8C=E5=8A=9F=E8=83=BD=E6=9C=AA=E7=BB=8F=E4=BB=BB=E4=BD=95?= =?UTF-8?q?=E6=B5=8B=E8=AF=95=EF=BC=8C=E4=BB=85=E4=B8=BA=E5=9F=BA=E6=9C=AC?= =?UTF-8?q?=E6=A1=86=E6=9E=B6=EF=BC=8C=E7=9B=AE=E5=89=8D=E4=B8=BA=E5=AE=8C?= =?UTF-8?q?=E5=85=A8=E4=B8=8D=E5=8F=AF=E7=94=A8=E7=8A=B6=E6=80=81?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- README.md | 197 ++++++++++++++++ backend/.env.example | 35 +++ backend/api/admin.go | 359 ++++++++++++++++++++++++++++++ backend/api/chat.go | 277 +++++++++++++++++++++++ backend/api/middleware.go | 98 ++++++++ backend/api/user.go | 274 +++++++++++++++++++++++ backend/config/config.go | 169 ++++++++++++++ backend/db/db.go | 127 +++++++++++ backend/go.mod | 49 ++++ backend/main.go | 87 ++++++++ backend/models/models.go | 106 +++++++++ backend/services/ai_service.go | 132 +++++++++++ backend/services/quota_service.go | 162 ++++++++++++++ backend/utils/auth.go | 92 ++++++++ env | 34 --- index.html | 79 +++++++ script.js | 157 +++++++++++++ styles.css | 329 +++++++++++++++++++++++++++ 18 files changed, 2729 insertions(+), 34 deletions(-) create mode 100644 README.md create mode 100644 backend/.env.example create mode 100644 backend/api/admin.go create mode 100644 backend/api/chat.go create mode 100644 backend/api/middleware.go create mode 100644 backend/api/user.go create mode 100644 backend/config/config.go create mode 100644 backend/db/db.go create mode 100644 backend/go.mod create mode 100644 backend/main.go create mode 100644 backend/models/models.go create mode 100644 backend/services/ai_service.go create mode 100644 backend/services/quota_service.go create mode 100644 backend/utils/auth.go delete mode 100644 env create mode 100644 index.html create mode 100644 script.js create mode 100644 styles.css diff --git a/README.md b/README.md new file mode 100644 index 0000000..3a6d67f --- /dev/null +++ b/README.md @@ -0,0 +1,197 @@ +# 豆包AI对话系统 + +一个模仿豆包官网的AI对话系统,包含前端页面和Go后端服务。 + +## 功能特性 + +### 前端功能 +- 📱 简洁美观的聊天界面 +- 💬 支持多轮对话 +- 📝 支持换行输入(Shift+Enter) +- ⌨️ 支持键盘快捷键 +- 📱 响应式设计,适配各种设备 + +### 后端功能 +- 🔐 用户认证和授权(JWT) +- 👥 多用户支持 +- 📊 用户配额管理 +- 🧠 调用自建大模型API +- 📝 对话历史记录 +- 👨💼 管理员管理功能 +- 📈 系统统计信息 + +## 技术栈 + +### 前端 +- HTML5 + CSS3 + JavaScript +- Font Awesome 图标库 + +### 后端 +- Go 1.21 +- Gin Web框架 +- GORM ORM框架 +- SQLite数据库(默认) +- JWT认证 + +## 目录结构 + +``` +dobaochet/ +├── index.html # 前端主页面 +├── styles.css # 前端样式 +├── script.js # 前端脚本 +├── backend/ # 后端代码 +│ ├── main.go # 主程序 +│ ├── go.mod # Go依赖 +│ ├── .env.example # 配置文件示例 +│ ├── api/ # API路由 +│ ├── models/ # 数据模型 +│ ├── services/ # 业务逻辑 +│ ├── config/ # 配置管理 +│ ├── utils/ # 工具函数 +│ └── db/ # 数据库连接 +└── README.md # 项目说明 +``` + +## 快速开始 + +### 1. 前端使用 + +直接在浏览器中打开 `index.html` 文件即可使用前端页面。 + +### 2. 后端部署 + +#### 2.1 配置环境 + +复制 `.env.example` 为 `.env` 并修改配置: + +```bash +cd backend +cp .env.example .env +``` + +编辑 `.env` 文件,配置以下关键参数: + +```env +# AI模型配置 +AI_API_URL=http://localhost:5000/v1/chat/completions # 替换为您的大模型API URL +AI_API_KEY= # 大模型API密钥 +AI_MODEL_NAME=default # 模型名称 +``` + +#### 2.2 安装依赖 + +```bash +go mod tidy +``` + +#### 2.3 启动服务 + +```bash +go run main.go +``` + +服务将在 `http://localhost:8000` 启动。 + +## API文档 + +### 认证相关 + +- `POST /api/users/register` - 用户注册 +- `POST /api/users/login` - 用户登录 +- `GET /api/users/me` - 获取当前用户信息 + +### 聊天相关 + +- `POST /api/chat/` - 发送聊天消息 +- `GET /api/chat/conversations` - 获取对话列表 +- `GET /api/chat/conversations/:id` - 获取对话详情 +- `DELETE /api/chat/conversations/:id` - 删除对话 +- `DELETE /api/chat/conversations` - 清空所有对话 + +### 管理员相关 + +- `GET /api/admin/users` - 获取所有用户 +- `GET /api/admin/users/:id` - 获取单个用户 +- `POST /api/admin/users` - 创建用户 +- `PUT /api/admin/users/:id` - 更新用户 +- `DELETE /api/admin/users/:id` - 删除用户 +- `PUT /api/admin/users/:id/quota` - 更新用户配额 +- `POST /api/admin/users/:id/quota/reset` - 重置用户配额 +- `GET /api/admin/stats` - 获取系统统计信息 + +## 配置说明 + +### 数据库配置 + +支持多种数据库,默认使用SQLite: + +```env +# SQLite +DATABASE_DSN=sqlite://./chat.db + +# MySQL +# DATABASE_DSN=mysql://user:password@tcp(localhost:3306)/chat?charset=utf8mb4&parseTime=True&loc=Local + +# PostgreSQL +# DATABASE_DSN=postgres://user:password@localhost:5432/chat?sslmode=disable +``` + +### JWT配置 + +```env +JWT_SECRET_KEY=your-secret-key-change-me # 生产环境请修改为复杂密钥 +JWT_ACCESS_TOKEN_EXPIRE=24h # 访问令牌过期时间 +JWT_REFRESH_TOKEN_EXPIRE=720h # 刷新令牌过期时间 +``` + +### 配额配置 + +```env +QUOTA_DEFAULT_TOTAL_TOKENS=100000 # 默认总配额 +QUOTA_DEFAULT_TOKEN_LIMIT=1000 # 单次请求限制 +QUOTA_RESET_INTERVAL=720h # 配额重置间隔 +``` + +## 默认管理员账号 + +系统初始化后会自动创建管理员账号: + +- 用户名:`admin` +- 密码:`admin123` +- 生产环境请及时修改密码! + +## 开发说明 + +### 前端开发 + +前端使用纯HTML+CSS+JavaScript开发,无需编译,直接修改文件即可。 + +### 后端开发 + +后端使用Go语言开发,推荐使用Go 1.21或以上版本。 + +## 生产环境部署 + +1. 修改JWT密钥为复杂随机字符串 +2. 配置正确的数据库连接 +3. 配置HTTPS(推荐使用Nginx反向代理) +4. 设置适当的日志级别 +5. 定期备份数据库 + +## 安全建议 + +1. 生产环境中不要使用默认的JWT密钥 +2. 定期更换JWT密钥 +3. 配置适当的CORS策略 +4. 启用HTTPS +5. 限制API请求频率 +6. 定期审计系统日志 + +## 许可证 + +MIT License + +## 贡献 + +欢迎提交Issue和Pull Request! diff --git a/backend/.env.example b/backend/.env.example new file mode 100644 index 0000000..4a5975d --- /dev/null +++ b/backend/.env.example @@ -0,0 +1,35 @@ +# 服务器配置 +SERVER_PORT=8000 +SERVER_HOST=0.0.0.0 +SERVER_READ_TIMEOUT=15s +SERVER_WRITE_TIMEOUT=15s + +# 数据库配置 +# 支持 SQLite, MySQL, PostgreSQL +DATABASE_DSN=sqlite://./chat.db +DATABASE_MAX_IDLE_CONNS=10 +DATABASE_MAX_OPEN_CONNS=100 + +# JWT 配置 +JWT_SECRET_KEY=your-secret-key-change-me +JWT_ACCESS_TOKEN_EXPIRE=24h +JWT_REFRESH_TOKEN_EXPIRE=720h + +# AI 模型配置 +# 替换为您自己的大模型 API URL +AI_API_URL=http://localhost:5000/v1/chat/completions +AI_API_KEY= +AI_MODEL_NAME=default +AI_MAX_TOKENS=2048 +AI_TEMPERATURE=0.7 + +# 配额配置 +QUOTA_DEFAULT_TOTAL_TOKENS=100000 +QUOTA_DEFAULT_TOKEN_LIMIT=1000 +QUOTA_RESET_INTERVAL=720h + +# CORS 配置 +CORS_ALLOW_ORIGINS=* +CORS_ALLOW_METHODS=GET,POST,PUT,DELETE,OPTIONS +CORS_ALLOW_HEADERS=Origin,Content-Type,Authorization +CORS_ALLOW_CREDENTIALS=true diff --git a/backend/api/admin.go b/backend/api/admin.go new file mode 100644 index 0000000..75d0729 --- /dev/null +++ b/backend/api/admin.go @@ -0,0 +1,359 @@ +package api + +import ( + "net/http" + "time" + + "d:\项目\dobaochet\backend\config" + "d:\项目\dobaochet\backend\db" + "d:\项目\dobaochet\backend\models" + "d:\项目\dobaochet\backend\utils" + + "github.com/gin-gonic/gin" +) + +// AdminAPI 管理员API结构体 +type AdminAPI struct { + cfg *config.Config +} + +// NewAdminAPI 创建管理员API实例 +func NewAdminAPI(cfg *config.Config) *AdminAPI { + return &AdminAPI{cfg: cfg} +} + +// GetUsers 获取所有用户 +func (api *AdminAPI) GetUsers(c *gin.Context) { + // 分页参数 + page := 1 + pageSize := 10 + c.Query("page") + c.Query("page_size") + + // 查询用户列表 + var users []models.User + var total int64 + + db.GetDB().Model(&models.User{}).Count(&total) + db.GetDB().Preload("Quota").Order("created_at desc").Offset((page - 1) * pageSize).Limit(pageSize).Find(&users) + + // 转换为响应格式 + var userResponses []models.UserResponse + for _, user := range users { + userResponses = append(userResponses, models.UserResponse{ + ID: user.ID, + Username: user.Username, + Email: user.Email, + IsAdmin: user.IsAdmin, + IsActive: user.IsActive, + CreatedAt: user.CreatedAt, + Quota: user.Quota, + }) + } + + c.JSON(http.StatusOK, gin.H{ + "users": userResponses, + "total": total, + "page": page, + "page_size": pageSize, + "total_pages": (total + int64(pageSize) - 1) / int64(pageSize), + }) +} + +// GetUser 获取单个用户 +func (api *AdminAPI) GetUser(c *gin.Context) { + // 获取用户ID + id := c.Param("id") + + // 查询用户 + var user models.User + if result := db.GetDB().Preload("Quota").First(&user, id); result.Error != nil { + c.JSON(http.StatusNotFound, gin.H{"error": "用户不存在"}) + return + } + + c.JSON(http.StatusOK, models.UserResponse{ + ID: user.ID, + Username: user.Username, + Email: user.Email, + IsAdmin: user.IsAdmin, + IsActive: user.IsActive, + CreatedAt: user.CreatedAt, + Quota: user.Quota, + }) +} + +// CreateUser 创建用户 +func (api *AdminAPI) CreateUser(c *gin.Context) { + var req struct { + Username string `json:"username" binding:"required,min=3,max=50"` + Email string `json:"email" binding:"required,email"` + Password string `json:"password" binding:"required,min=6"` + IsAdmin bool `json:"is_admin"` + IsActive bool `json:"is_active"` + } + + if err := c.ShouldBindJSON(&req); err != nil { + c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) + return + } + + // 检查用户名是否已存在 + var existingUser models.User + if result := db.GetDB().Where("username = ?", req.Username).First(&existingUser); result.Error == nil { + c.JSON(http.StatusConflict, gin.H{"error": "用户名已存在"}) + return + } + + // 检查邮箱是否已存在 + if result := db.GetDB().Where("email = ?", req.Email).First(&existingUser); result.Error == nil { + c.JSON(http.StatusConflict, gin.H{"error": "邮箱已存在"}) + return + } + + // 密码加密 + hashedPassword, err := utils.HashPassword(req.Password) + if err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": "密码加密失败"}) + return + } + + // 创建用户 + user := models.User{ + Username: req.Username, + Email: req.Email, + Password: hashedPassword, + IsAdmin: req.IsAdmin, + IsActive: req.IsActive, + LastLoginAt: time.Now(), + Quota: models.Quota{ + TotalTokens: api.cfg.Quota.DefaultTotalTokens, + UsedTokens: 0, + ResetAt: time.Now().Add(api.cfg.Quota.ResetInterval), + TokenLimit: api.cfg.Quota.DefaultTokenLimit, + }, + } + + if result := db.GetDB().Create(&user); result.Error != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": "创建用户失败"}) + return + } + + c.JSON(http.StatusCreated, models.UserResponse{ + ID: user.ID, + Username: user.Username, + Email: user.Email, + IsAdmin: user.IsAdmin, + IsActive: user.IsActive, + CreatedAt: user.CreatedAt, + Quota: user.Quota, + }) +} + +// UpdateUser 更新用户 +func (api *AdminAPI) UpdateUser(c *gin.Context) { + // 获取用户ID + id := c.Param("id") + + // 查询用户 + var user models.User + if result := db.GetDB().First(&user, id); result.Error != nil { + c.JSON(http.StatusNotFound, gin.H{"error": "用户不存在"}) + return + } + + // 解析请求 + var req struct { + Email string `json:"email" binding:"omitempty,email"` + IsAdmin bool `json:"is_admin"` + IsActive bool `json:"is_active"` + } + + if err := c.ShouldBindJSON(&req); err != nil { + c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) + return + } + + // 更新用户信息 + if req.Email != "" && req.Email != user.Email { + // 检查邮箱是否已被使用 + var existingUser models.User + if result := db.GetDB().Where("email = ? AND id != ?", req.Email, user.ID).First(&existingUser); result.Error == nil { + c.JSON(http.StatusConflict, gin.H{"error": "邮箱已被使用"}) + return + } + user.Email = req.Email + } + + user.IsAdmin = req.IsAdmin + user.IsActive = req.IsActive + + if result := db.GetDB().Save(&user); result.Error != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": "更新用户失败"}) + return + } + + // 重新加载用户信息 + db.GetDB().Preload("Quota").First(&user, id) + + c.JSON(http.StatusOK, models.UserResponse{ + ID: user.ID, + Username: user.Username, + Email: user.Email, + IsAdmin: user.IsAdmin, + IsActive: user.IsActive, + CreatedAt: user.CreatedAt, + Quota: user.Quota, + }) +} + +// DeleteUser 删除用户 +func (api *AdminAPI) DeleteUser(c *gin.Context) { + // 获取用户ID + id := c.Param("id") + + // 检查是否是最后一个管理员 + var user models.User + if result := db.GetDB().First(&user, id); result.Error != nil { + c.JSON(http.StatusNotFound, gin.H{"error": "用户不存在"}) + return + } + + if user.IsAdmin { + var adminCount int64 + db.GetDB().Model(&models.User{}).Where("is_admin = ?", true).Count(&adminCount) + if adminCount <= 1 { + c.JSON(http.StatusBadRequest, gin.H{"error": "不能删除最后一个管理员用户"}) + return + } + } + + // 删除用户(级联删除相关数据) + if result := db.GetDB().Delete(&user); result.Error != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": "删除用户失败"}) + return + } + + c.JSON(http.StatusOK, gin.H{"message": "用户删除成功"}) +} + +// UpdateUserQuota 更新用户配额 +func (api *AdminAPI) UpdateUserQuota(c *gin.Context) { + // 获取用户ID + id := c.Param("id") + + // 查询用户 + var user models.User + if result := db.GetDB().Preload("Quota").First(&user, id); result.Error != nil { + c.JSON(http.StatusNotFound, gin.H{"error": "用户不存在"}) + return + } + + // 解析请求 + var req models.QuotaUpdateRequest + if err := c.ShouldBindJSON(&req); err != nil { + c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) + return + } + + // 更新配额 + user.Quota.TotalTokens = req.TotalTokens + user.Quota.TokenLimit = req.TokenLimit + + if result := db.GetDB().Save(&user.Quota); result.Error != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": "更新配额失败"}) + return + } + + c.JSON(http.StatusOK, gin.H{ + "message": "配额更新成功", + "quota": user.Quota, + }) +} + +// ResetUserQuota 重置用户配额 +func (api *AdminAPI) ResetUserQuota(c *gin.Context) { + // 获取用户ID + id := c.Param("id") + + // 查询用户 + var user models.User + if result := db.GetDB().Preload("Quota").First(&user, id); result.Error != nil { + c.JSON(http.StatusNotFound, gin.H{"error": "用户不存在"}) + return + } + + // 重置配额 + user.Quota.UsedTokens = 0 + user.Quota.ResetAt = time.Now().Add(api.cfg.Quota.ResetInterval) + + if result := db.GetDB().Save(&user.Quota); result.Error != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": "重置配额失败"}) + return + } + + c.JSON(http.StatusOK, gin.H{ + "message": "配额重置成功", + "quota": user.Quota, + }) +} + +// GetSystemStats 获取系统统计信息 +func (api *AdminAPI) GetSystemStats(c *gin.Context) { + // 用户统计 + var totalUsers, activeUsers, adminUsers int64 + db.GetDB().Model(&models.User{}).Count(&totalUsers) + db.GetDB().Model(&models.User{}).Where("is_active = ?", true).Count(&activeUsers) + db.GetDB().Model(&models.User{}).Where("is_admin = ?", true).Count(&adminUsers) + + // 对话统计 + var totalConversations int64 + db.GetDB().Model(&models.Conversation{}).Count(&totalConversations) + + // 消息统计 + var totalMessages int64 + db.GetDB().Model(&models.Message{}).Count(&totalMessages) + + // 配额统计 + var totalTokens, usedTokens int64 + db.GetDB().Model(&models.Quota{}).Select("sum(total_tokens) as total_tokens, sum(used_tokens) as used_tokens").Scan(&struct { + TotalTokens int64 + UsedTokens int64 + }{TotalTokens: &totalTokens, UsedTokens: &usedTokens}) + + c.JSON(http.StatusOK, gin.H{ + "users": gin.H{ + "total": totalUsers, + "active": activeUsers, + "admin": adminUsers, + }, + "conversations": totalConversations, + "messages": totalMessages, + "quota": gin.H{ + "total_tokens": totalTokens, + "used_tokens": usedTokens, + "usage_rate": float64(usedTokens) / float64(totalTokens) * 100, + }, + }) +} + +// RegisterRoutes 注册管理员路由 +func (api *AdminAPI) RegisterRoutes(router *gin.RouterGroup) { + adminGroup := router.Group("/admin") + adminGroup.Use(AuthMiddleware(api.cfg), AdminMiddleware()) + { + // 用户管理 + adminGroup.GET("/users", api.GetUsers) + adminGroup.GET("/users/:id", api.GetUser) + adminGroup.POST("/users", api.CreateUser) + adminGroup.PUT("/users/:id", api.UpdateUser) + adminGroup.DELETE("/users/:id", api.DeleteUser) + + // 配额管理 + adminGroup.PUT("/users/:id/quota", api.UpdateUserQuota) + adminGroup.POST("/users/:id/quota/reset", api.ResetUserQuota) + + // 系统统计 + adminGroup.GET("/stats", api.GetSystemStats) + } +} diff --git a/backend/api/chat.go b/backend/api/chat.go new file mode 100644 index 0000000..bb3438c --- /dev/null +++ b/backend/api/chat.go @@ -0,0 +1,277 @@ +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 +} diff --git a/backend/api/middleware.go b/backend/api/middleware.go new file mode 100644 index 0000000..496005f --- /dev/null +++ b/backend/api/middleware.go @@ -0,0 +1,98 @@ +package api + +import ( + "net/http" + "strings" + + "d:\项目\dobaochet\backend\config" + "d:\项目\dobaochet\backend\db" + "d:\项目\dobaochet\backend\models" + "d:\项目\dobaochet\backend\utils" + + "github.com/gin-gonic/gin" +) + +// AuthMiddleware 认证中间件 +func AuthMiddleware(cfg *config.Config) gin.HandlerFunc { + return func(c *gin.Context) { + // 从请求头获取Authorization + authHeader := c.GetHeader("Authorization") + if authHeader == "" { + c.JSON(http.StatusUnauthorized, gin.H{"error": "未提供认证令牌"}) + c.Abort() + return + } + + // 检查Bearer前缀 + parts := strings.SplitN(authHeader, " ", 2) + if !(len(parts) == 2 && parts[0] == "Bearer") { + c.JSON(http.StatusUnauthorized, gin.H{"error": "认证令牌格式错误"}) + c.Abort() + return + } + + // 解析JWT令牌 + claims, err := utils.ParseToken(parts[1], cfg.JWT.SecretKey) + if err != nil { + c.JSON(http.StatusUnauthorized, gin.H{"error": "无效的认证令牌"}) + c.Abort() + return + } + + // 查询用户信息 + var user models.User + result := db.GetDB().Preload("Quota").First(&user, claims.UserID) + if result.Error != nil { + c.JSON(http.StatusUnauthorized, gin.H{"error": "用户不存在"}) + c.Abort() + return + } + + // 检查用户是否激活 + if !user.IsActive { + c.JSON(http.StatusForbidden, gin.H{"error": "用户已被禁用"}) + c.Abort() + return + } + + // 将用户信息设置到上下文 + c.Set("user", user) + c.Set("user_id", user.ID) + c.Set("is_admin", user.IsAdmin) + + c.Next() + } +} + +// AdminMiddleware 管理员中间件 +func AdminMiddleware() gin.HandlerFunc { + return func(c *gin.Context) { + // 从上下文获取用户信息 + isAdmin, exists := c.Get("is_admin") + if !exists || !isAdmin.(bool) { + c.JSON(http.StatusForbidden, gin.H{"error": "需要管理员权限"}) + c.Abort() + return + } + + c.Next() + } +} + +// CORSMiddleware CORS中间件 +func CORSMiddleware(cfg *config.Config) gin.HandlerFunc { + return func(c *gin.Context) { + // 设置允许的来源 + c.Writer.Header().Set("Access-Control-Allow-Origin", strings.Join(cfg.CORS.AllowOrigins, ",")) + c.Writer.Header().Set("Access-Control-Allow-Credentials", "true") + c.Writer.Header().Set("Access-Control-Allow-Headers", strings.Join(cfg.CORS.AllowHeaders, ",")) + c.Writer.Header().Set("Access-Control-Allow-Methods", strings.Join(cfg.CORS.AllowMethods, ",")) + + if c.Request.Method == "OPTIONS" { + c.AbortWithStatus(204) + return + } + + c.Next() + } +} diff --git a/backend/api/user.go b/backend/api/user.go new file mode 100644 index 0000000..c6655ff --- /dev/null +++ b/backend/api/user.go @@ -0,0 +1,274 @@ +package api + +import ( + "net/http" + "time" + + "d:\项目\dobaochet\backend\config" + "d:\项目\dobaochet\backend\db" + "d:\项目\dobaochet\backend\models" + "d:\项目\dobaochet\backend\utils" + + "github.com/gin-gonic/gin" +) + +// UserAPI 用户API结构体 +type UserAPI struct { + cfg *config.Config +} + +// NewUserAPI 创建用户API实例 +func NewUserAPI(cfg *config.Config) *UserAPI { + return &UserAPI{cfg: cfg} +} + +// Register 注册用户 +func (api *UserAPI) Register(c *gin.Context) { + var req models.RegisterRequest + if err := c.ShouldBindJSON(&req); err != nil { + c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) + return + } + + // 检查用户名是否已存在 + var existingUser models.User + if result := db.GetDB().Where("username = ?", req.Username).First(&existingUser); result.Error == nil { + c.JSON(http.StatusConflict, gin.H{"error": "用户名已存在"}) + return + } + + // 检查邮箱是否已存在 + if result := db.GetDB().Where("email = ?", req.Email).First(&existingUser); result.Error == nil { + c.JSON(http.StatusConflict, gin.H{"error": "邮箱已存在"}) + return + } + + // 密码加密 + hashedPassword, err := utils.HashPassword(req.Password) + if err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": "密码加密失败"}) + return + } + + // 创建用户 + user := models.User{ + Username: req.Username, + Email: req.Email, + Password: hashedPassword, + IsActive: true, + LastLoginAt: time.Now(), + Quota: models.Quota{ + TotalTokens: api.cfg.Quota.DefaultTotalTokens, + UsedTokens: 0, + ResetAt: time.Now().Add(api.cfg.Quota.ResetInterval), + TokenLimit: api.cfg.Quota.DefaultTokenLimit, + }, + } + + if result := db.GetDB().Create(&user); result.Error != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": "创建用户失败"}) + return + } + + // 返回用户信息(隐藏敏感信息) + c.JSON(http.StatusCreated, gin.H{ + "id": user.ID, + "username": user.Username, + "email": user.Email, + "is_admin": user.IsAdmin, + "is_active": user.IsActive, + "created_at": user.CreatedAt, + }) +} + +// Login 用户登录 +func (api *UserAPI) Login(c *gin.Context) { + var req models.LoginRequest + if err := c.ShouldBindJSON(&req); err != nil { + c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) + return + } + + // 查找用户 + var user models.User + if result := db.GetDB().Preload("Quota").Where("username = ?", req.Username).First(&user); result.Error != nil { + c.JSON(http.StatusUnauthorized, gin.H{"error": "用户名或密码错误"}) + return + } + + // 验证密码 + if err := utils.VerifyPassword(user.Password, req.Password); err != nil { + c.JSON(http.StatusUnauthorized, gin.H{"error": "用户名或密码错误"}) + return + } + + // 检查用户是否激活 + if !user.IsActive { + c.JSON(http.StatusForbidden, gin.H{"error": "用户已被禁用"}) + return + } + + // 更新登录时间 + user.LastLoginAt = time.Now() + db.GetDB().Save(&user) + + // 生成JWT令牌 + token, err := utils.GenerateToken( + user.ID, + user.Username, + user.IsAdmin, + api.cfg.JWT.SecretKey, + api.cfg.JWT.AccessTokenExpire, + ) + if err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": "生成令牌失败"}) + return + } + + // 返回登录结果 + c.JSON(http.StatusOK, models.LoginResponse{ + Token: token, + User: models.UserResponse{ + ID: user.ID, + Username: user.Username, + Email: user.Email, + IsAdmin: user.IsAdmin, + IsActive: user.IsActive, + CreatedAt: user.CreatedAt, + Quota: user.Quota, + }, + }) +} + +// GetUserInfo 获取当前用户信息 +func (api *UserAPI) GetUserInfo(c *gin.Context) { + // 从上下文获取用户信息 + user, exists := c.Get("user") + if !exists { + c.JSON(http.StatusUnauthorized, gin.H{"error": "未认证"}) + return + } + + currentUser := user.(models.User) + c.JSON(http.StatusOK, models.UserResponse{ + ID: currentUser.ID, + Username: currentUser.Username, + Email: currentUser.Email, + IsAdmin: currentUser.IsAdmin, + IsActive: currentUser.IsActive, + CreatedAt: currentUser.CreatedAt, + Quota: currentUser.Quota, + }) +} + +// UpdateUserInfo 更新用户信息 +func (api *UserAPI) UpdateUserInfo(c *gin.Context) { + // 从上下文获取用户信息 + user, exists := c.Get("user") + if !exists { + c.JSON(http.StatusUnauthorized, gin.H{"error": "未认证"}) + return + } + + currentUser := user.(models.User) + + // 解析请求 + var req struct { + Email string `json:"email" binding:"omitempty,email"` + } + + if err := c.ShouldBindJSON(&req); err != nil { + c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) + return + } + + // 更新用户信息 + if req.Email != "" && req.Email != currentUser.Email { + // 检查邮箱是否已被使用 + var existingUser models.User + if result := db.GetDB().Where("email = ? AND id != ?", req.Email, currentUser.ID).First(&existingUser); result.Error == nil { + c.JSON(http.StatusConflict, gin.H{"error": "邮箱已被使用"}) + return + } + currentUser.Email = req.Email + } + + if result := db.GetDB().Save(¤tUser); result.Error != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": "更新用户信息失败"}) + return + } + + c.JSON(http.StatusOK, models.UserResponse{ + ID: currentUser.ID, + Username: currentUser.Username, + Email: currentUser.Email, + IsAdmin: currentUser.IsAdmin, + IsActive: currentUser.IsActive, + CreatedAt: currentUser.CreatedAt, + Quota: currentUser.Quota, + }) +} + +// ChangePassword 修改密码 +func (api *UserAPI) ChangePassword(c *gin.Context) { + // 从上下文获取用户信息 + user, exists := c.Get("user") + if !exists { + c.JSON(http.StatusUnauthorized, gin.H{"error": "未认证"}) + return + } + + currentUser := user.(models.User) + + // 解析请求 + var req struct { + OldPassword string `json:"old_password" binding:"required"` + NewPassword string `json:"new_password" binding:"required,min=6"` + } + + if err := c.ShouldBindJSON(&req); err != nil { + c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) + return + } + + // 验证旧密码 + if err := utils.VerifyPassword(currentUser.Password, req.OldPassword); err != nil { + c.JSON(http.StatusBadRequest, gin.H{"error": "旧密码错误"}) + return + } + + // 加密新密码 + newHashedPassword, err := utils.HashPassword(req.NewPassword) + if err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": "密码加密失败"}) + return + } + + // 更新密码 + currentUser.Password = newHashedPassword + if result := db.GetDB().Save(¤tUser); result.Error != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": "修改密码失败"}) + return + } + + c.JSON(http.StatusOK, gin.H{"message": "密码修改成功"}) +} + +// RegisterRoutes 注册用户路由 +func (api *UserAPI) RegisterRoutes(router *gin.RouterGroup) { + userGroup := router.Group("/users") + { + // 公开路由 + userGroup.POST("/register", api.Register) + userGroup.POST("/login", api.Login) + + // 需要认证的路由 + authGroup := userGroup.Group("/") + authGroup.Use(AuthMiddleware(api.cfg)) + { + authGroup.GET("/me", api.GetUserInfo) + authGroup.PUT("/me", api.UpdateUserInfo) + authGroup.PUT("/me/password", api.ChangePassword) + } + } +} diff --git a/backend/config/config.go b/backend/config/config.go new file mode 100644 index 0000000..b86ed10 --- /dev/null +++ b/backend/config/config.go @@ -0,0 +1,169 @@ +package config + +import ( + "fmt" + "log" + "os" + "time" + + "github.com/spf13/viper" +) + +// Config 应用配置结构体 +type Config struct { + Server ServerConfig `mapstructure:"server"` + Database DatabaseConfig `mapstructure:"database"` + JWT JWTConfig `mapstructure:"jwt"` + AI AIConfig `mapstructure:"ai"` + Quota QuotaConfig `mapstructure:"quota"` + CORS CORSConfig `mapstructure:"cors"` +} + +// ServerConfig 服务器配置 +type ServerConfig struct { + Port string `mapstructure:"port"` + Host string `mapstructure:"host"` + ReadTimeout time.Duration `mapstructure:"read_timeout"` + WriteTimeout time.Duration `mapstructure:"write_timeout"` +} + +// DatabaseConfig 数据库配置 +type DatabaseConfig struct { + DSN string `mapstructure:"dsn"` + MaxIdleConns int `mapstructure:"max_idle_conns"` + MaxOpenConns int `mapstructure:"max_open_conns"` +} + +// JWTConfig JWT配置 +type JWTConfig struct { + SecretKey string `mapstructure:"secret_key"` + AccessTokenExpire time.Duration `mapstructure:"access_token_expire"` + RefreshTokenExpire time.Duration `mapstructure:"refresh_token_expire"` +} + +// AIConfig AI模型配置 +type AIConfig struct { + APIURL string `mapstructure:"api_url"` + APIKey string `mapstructure:"api_key"` + ModelName string `mapstructure:"model_name"` + MaxTokens int64 `mapstructure:"max_tokens"` + Temperature float64 `mapstructure:"temperature"` +} + +// QuotaConfig 配额配置 +type QuotaConfig struct { + DefaultTotalTokens int64 `mapstructure:"default_total_tokens"` + DefaultTokenLimit int64 `mapstructure:"default_token_limit"` + ResetInterval time.Duration `mapstructure:"reset_interval"` +} + +// CORSConfig CORS配置 +type CORSConfig struct { + AllowOrigins []string `mapstructure:"allow_origins"` + AllowMethods []string `mapstructure:"allow_methods"` + AllowHeaders []string `mapstructure:"allow_headers"` + AllowCredentials bool `mapstructure:"allow_credentials"` +} + +// LoadConfig 加载配置 +func LoadConfig(path string) (*Config, error) { + // 设置默认值 + setDefaults() + + // 读取配置文件 + viper.AddConfigPath(path) + viper.SetConfigName(".env") + viper.SetConfigType("env") + + // 读取环境变量 + viper.AutomaticEnv() + + // 尝试读取配置文件 + if err := viper.ReadInConfig(); err != nil { + if _, ok := err.(viper.ConfigFileNotFoundError); ok { + log.Println("警告: 未找到配置文件,将使用环境变量和默认值") + } else { + return nil, fmt.Errorf("读取配置文件错误: %w", err) + } + } + + // 解析配置 + var config Config + if err := viper.Unmarshal(&config); err != nil { + return nil, fmt.Errorf("解析配置错误: %w", err) + } + + // 验证必要的配置 + if err := validateConfig(&config); err != nil { + return nil, err + } + + return &config, nil +} + +// 设置默认值 +func setDefaults() { + // 服务器默认配置 + viper.SetDefault("server.port", "8000") + viper.SetDefault("server.host", "0.0.0.0") + viper.SetDefault("server.read_timeout", 15*time.Second) + viper.SetDefault("server.write_timeout", 15*time.Second) + + // 数据库默认配置(使用SQLite) + viper.SetDefault("database.dsn", "./chat.db") + viper.SetDefault("database.max_idle_conns", 10) + viper.SetDefault("database.max_open_conns", 100) + + // JWT默认配置 + viper.SetDefault("jwt.secret_key", "your-secret-key-change-me") + viper.SetDefault("jwt.access_token_expire", 24*time.Hour) + viper.SetDefault("jwt.refresh_token_expire", 7*24*time.Hour) + + // AI模型默认配置 + viper.SetDefault("ai.api_url", "http://localhost:5000/v1/chat/completions") + viper.SetDefault("ai.api_key", "") + viper.SetDefault("ai.model_name", "default") + viper.SetDefault("ai.max_tokens", 2048) + viper.SetDefault("ai.temperature", 0.7) + + // 配额默认配置 + viper.SetDefault("quota.default_total_tokens", 100000) + viper.SetDefault("quota.default_token_limit", 1000) + viper.SetDefault("quota.reset_interval", 30*24*time.Hour) + + // CORS默认配置 + viper.SetDefault("cors.allow_origins", []string{"*"}) + viper.SetDefault("cors.allow_methods", []string{"GET", "POST", "PUT", "DELETE", "OPTIONS"}) + viper.SetDefault("cors.allow_headers", []string{"Origin", "Content-Type", "Authorization"}) + viper.SetDefault("cors.allow_credentials", true) +} + +// 验证配置 +func validateConfig(config *Config) error { + // 验证JWT密钥 + if config.JWT.SecretKey == "your-secret-key-change-me" { + log.Println("警告: 使用默认JWT密钥,建议在生产环境中修改") + } + + // 验证AI API URL + if config.AI.APIURL == "" { + return fmt.Errorf("AI API URL不能为空") + } + + return nil +} + +// GetDSN 获取数据库DSN +func (c *DatabaseConfig) GetDSN() string { + // 如果是SQLite,确保目录存在 + if len(c.DSN) > 6 && c.DSN[:6] == "sqlite" { + // 提取SQLite文件路径 + filePath := c.DSN[9:] // 去掉 "sqlite://" 前缀 + dir := filePath[:len(filePath)-len("/"+filePath[len(filePath)-1:])] + if dir != "" { + os.MkdirAll(dir, 0755) + } + } + + return c.DSN +} diff --git a/backend/db/db.go b/backend/db/db.go new file mode 100644 index 0000000..a36d911 --- /dev/null +++ b/backend/db/db.go @@ -0,0 +1,127 @@ +package db + +import ( + "log" + "time" + + "d:\项目\dobaochet\backend\config" + "d:\项目\dobaochet\backend\models" + "d:\项目\dobaochet\backend\utils" + + "gorm.io/driver/sqlite" + "gorm.io/gorm" + "gorm.io/gorm/logger" +) + +// DB 全局数据库连接 +var DB *gorm.DB + +// InitDB 初始化数据库连接 +func InitDB(cfg *config.Config) (*gorm.DB, error) { + // 配置GORM日志 + newLogger := logger.New( + log.New(log.Writer(), "\r\n", log.LstdFlags), + logger.Config{ + SlowThreshold: time.Second, // 慢SQL阈值 + LogLevel: logger.Info, // 日志级别 + Colorful: true, // 彩色日志 + }, + ) + + // 连接数据库 + var err error + switch { + case len(cfg.Database.DSN) > 6 && cfg.Database.DSN[:6] == "sqlite": + // SQLite连接 + DB, err = gorm.Open(sqlite.Open(cfg.Database.DSN[9:]), &gorm.Config{ + Logger: newLogger, + }) + default: + // 默认使用SQLite + DB, err = gorm.Open(sqlite.Open("./chat.db"), &gorm.Config{ + Logger: newLogger, + }) + } + + if err != nil { + return nil, err + } + + // 设置连接池 + sqlDB, err := DB.DB() + if err != nil { + return nil, err + } + + sqlDB.SetMaxIdleConns(cfg.Database.MaxIdleConns) + sqlDB.SetMaxOpenConns(cfg.Database.MaxOpenConns) + sqlDB.SetConnMaxLifetime(time.Hour) + + // 自动迁移模型 + err = migrateModels() + if err != nil { + return nil, err + } + + // 初始化数据 + err = seedData(cfg) + if err != nil { + return nil, err + } + + return DB, nil +} + +// migrateModels 自动迁移模型 +func migrateModels() error { + return DB.AutoMigrate( + &models.User{}, + &models.Quota{}, + &models.Conversation{}, + &models.Message{}, + ) +} + +// seedData 初始化数据 +func seedData(cfg *config.Config) error { + // 检查是否已有管理员用户 + var count int64 + DB.Model(&models.User{}).Where("is_admin = ?", true).Count(&count) + + if count == 0 { + // 创建管理员用户 + password, err := utils.HashPassword("admin123") + if err != nil { + return err + } + + adminUser := models.User{ + Username: "admin", + Email: "admin@example.com", + Password: password, + IsAdmin: true, + IsActive: true, + LastLoginAt: time.Now(), + Quota: models.Quota{ + TotalTokens: cfg.Quota.DefaultTotalTokens * 10, // 管理员配额10倍 + UsedTokens: 0, + ResetAt: time.Now().Add(cfg.Quota.ResetInterval), + TokenLimit: cfg.Quota.DefaultTokenLimit * 10, + }, + } + + result := DB.Create(&adminUser) + if result.Error != nil { + return result.Error + } + + log.Println("已创建管理员用户: admin / admin123") + } + + return nil +} + +// GetDB 获取数据库连接 +func GetDB() *gorm.DB { + return DB +} diff --git a/backend/go.mod b/backend/go.mod new file mode 100644 index 0000000..54164b9 --- /dev/null +++ b/backend/go.mod @@ -0,0 +1,49 @@ +module dobaochet + +go 1.21 + +require ( + github.com/gin-gonic/gin v1.9.1 + github.com/golang-jwt/jwt/v5 v5.2.0 + github.com/spf13/viper v1.18.2 + golang.org/x/crypto v0.17.0 + gorm.io/driver/sqlite v1.5.4 + gorm.io/gorm v1.25.4 +) + +require ( + github.com/bytedance/sonic v1.10.1 // indirect + github.com/chenzhuoyu/base64x v0.0.0-20230717121745-296ad89f973d // indirect + github.com/fsnotify/fsnotify v1.7.0 // indirect + github.com/gabriel-vasile/mimetype v1.4.2 // indirect + github.com/gin-contrib/sse v0.1.0 // indirect + github.com/go-playground/locales v0.14.1 // indirect + github.com/go-playground/universal-translator v0.18.1 // indirect + github.com/go-playground/validator/v10 v10.16.0 // indirect + github.com/goccy/go-json v0.10.2 // indirect + github.com/hashicorp/hcl v1.0.0 // indirect + github.com/jinzhu/inflection v1.0.0 // indirect + github.com/jinzhu/now v1.1.5 // indirect + github.com/json-iterator/go v1.1.12 // indirect + github.com/klauspost/cpuid/v2 v2.2.5 // indirect + github.com/leodido/go-urn v1.2.4 // indirect + github.com/magiconair/properties v1.8.7 // indirect + github.com/mattn/go-isatty v0.0.20 // indirect + github.com/mitchellh/mapstructure v1.5.0 // indirect + github.com/modern-go/concurrent v0.0.0-20180306012644-bacd9c7ef1dd // indirect + github.com/modern-go/reflect2 v1.0.2 // indirect + github.com/pelletier/go-toml/v2 v2.1.0 // indirect + github.com/spf13/afero v1.11.0 // indirect + github.com/spf13/cast v1.6.0 // indirect + github.com/spf13/jwalterweatherman v1.1.0 // indirect + github.com/spf13/pflag v1.0.5 // indirect + github.com/subosito/gotenv v1.6.0 // indirect + github.com/twitchyliquid64/golang-asm v0.15.1 // indirect + github.com/ugorji/go/codec v1.2.11 // indirect + golang.org/x/arch v0.6.0 // indirect + golang.org/x/net v0.19.0 // indirect + golang.org/x/sys v0.15.0 // indirect + golang.org/x/text v0.14.0 // indirect + gopkg.in/ini.v1 v1.67.0 // indirect + gopkg.in/yaml.v3 v3.0.1 // indirect +) diff --git a/backend/main.go b/backend/main.go new file mode 100644 index 0000000..f1935a8 --- /dev/null +++ b/backend/main.go @@ -0,0 +1,87 @@ +package main + +import ( + "fmt" + "log" + "net/http" + + "d:\项目\dobaochet\backend\api" + "d:\项目\dobaochet\backend\config" + "d:\项目\dobaochet\backend\db" + "d:\项目\dobaochet\backend\services" + + "github.com/gin-gonic/gin" +) + +func main() { + // 加载配置 + cfg, err := config.LoadConfig(".") + if err != nil { + log.Fatalf("加载配置失败: %v", err) + } + + // 初始化数据库 + _, err = db.InitDB(cfg) + if err != nil { + log.Fatalf("初始化数据库失败: %v", err) + } + + // 创建Gin引擎 + gin.SetMode(gin.ReleaseMode) + r := gin.New() + + // 添加中间件 + r.Use(gin.Logger()) + r.Use(gin.Recovery()) + r.Use(api.CORSMiddleware(cfg)) + + // 创建服务实例 + aiService := services.NewAIService(cfg) + quotaService := services.NewQuotaService(cfg) + + // 启动配额重置定时任务 + quotaService.StartQuotaResetJob() + + // 创建API实例 + userAPI := api.NewUserAPI(cfg) + chatAPI := api.NewChatAPI(cfg, aiService) + adminAPI := api.NewAdminAPI(cfg) + + // 注册路由 + apiGroup := r.Group("/api") + { + // 用户相关路由 + userAPI.RegisterRoutes(apiGroup) + + // 聊天相关路由 + chatAPI.RegisterRoutes(apiGroup) + + // 管理员相关路由 + adminAPI.RegisterRoutes(apiGroup) + + // 健康检查 + apiGroup.GET("/health", func(c *gin.Context) { + c.JSON(http.StatusOK, gin.H{ + "status": "ok", + "message": "豆包AI后端服务运行正常", + }) + }) + } + + // 启动服务器 + serverAddr := fmt.Sprintf("%s:%s", cfg.Server.Host, cfg.Server.Port) + log.Printf("服务器启动在 %s", serverAddr) + + // 创建HTTP服务器 + server := &http.Server{ + Addr: serverAddr, + Handler: r, + ReadTimeout: cfg.Server.ReadTimeout, + WriteTimeout: cfg.Server.WriteTimeout, + } + + // 启动服务器 + if err := server.ListenAndServe(); err != nil && err != http.ErrServerClosed { + log.Fatalf("服务器启动失败: %v", err) + } +} diff --git a/backend/models/models.go b/backend/models/models.go new file mode 100644 index 0000000..9da165a --- /dev/null +++ b/backend/models/models.go @@ -0,0 +1,106 @@ +package models + +import ( + "time" + + "gorm.io/gorm" +) + +// BaseModel 基础模型,包含通用字段 +type BaseModel struct { + ID uint `gorm:"primaryKey" json:"id"` + CreatedAt time.Time `json:"created_at"` + UpdatedAt time.Time `json:"updated_at"` + DeletedAt gorm.DeletedAt `gorm:"index" json:"deleted_at,omitempty"` +} + +// User 用户模型 +type User struct { + BaseModel + Username string `gorm:"uniqueIndex;size:50;not null" json:"username"` + Password string `gorm:"size:255;not null" json:"-"` + Email string `gorm:"uniqueIndex;size:100;not null" json:"email"` + IsAdmin bool `gorm:"default:false" json:"is_admin"` + IsActive bool `gorm:"default:true" json:"is_active"` + LastLoginAt time.Time `json:"last_login_at"` + Quota Quota `gorm:"foreignKey:UserID" json:"quota"` + Conversations []Conversation `gorm:"foreignKey:UserID" json:"conversations,omitempty"` +} + +// Quota 配额模型 +type Quota struct { + BaseModel + UserID uint `gorm:"uniqueIndex;not null" json:"user_id"` + TotalTokens int64 `gorm:"default:100000" json:"total_tokens"` + UsedTokens int64 `gorm:"default:0" json:"used_tokens"` + ResetAt time.Time `json:"reset_at"` + TokenLimit int64 `gorm:"default:1000" json:"token_limit"` // 单次请求限制 +} + +// Conversation 对话模型 +type Conversation struct { + BaseModel + UserID uint `gorm:"index;not null" json:"user_id"` + Title string `gorm:"size:200;default:'新对话'" json:"title"` + Messages []Message `gorm:"foreignKey:ConversationID" json:"messages,omitempty"` + TotalTokens int64 `gorm:"default:0" json:"total_tokens"` +} + +// Message 消息模型 +type Message struct { + BaseModel + ConversationID uint `gorm:"index;not null" json:"conversation_id"` + Role string `gorm:"size:20;not null" json:"role"` // user, assistant + Content string `gorm:"type:text;not null" json:"content"` + Tokens int64 `gorm:"default:0" json:"tokens"` + Model string `gorm:"size:50;default:'default'" json:"model"` +} + +// ChatRequest 聊天请求模型 +type ChatRequest struct { + Message string `json:"message" binding:"required"` + ConversationID uint `json:"conversation_id,omitempty"` +} + +// ChatResponse 聊天响应模型 +type ChatResponse struct { + Response string `json:"response"` + ConversationID uint `json:"conversation_id"` + TokensUsed int64 `json:"tokens_used"` +} + +// UserResponse 用户响应模型(隐藏敏感信息) +type UserResponse struct { + ID uint `json:"id"` + Username string `json:"username"` + Email string `json:"email"` + IsAdmin bool `json:"is_admin"` + IsActive bool `json:"is_active"` + CreatedAt time.Time `json:"created_at"` + Quota Quota `json:"quota"` +} + +// LoginRequest 登录请求模型 +type LoginRequest struct { + Username string `json:"username" binding:"required"` + Password string `json:"password" binding:"required"` +} + +// LoginResponse 登录响应模型 +type LoginResponse struct { + Token string `json:"token"` + User UserResponse `json:"user"` +} + +// RegisterRequest 注册请求模型 +type RegisterRequest struct { + Username string `json:"username" binding:"required,min=3,max=50"` + Email string `json:"email" binding:"required,email"` + Password string `json:"password" binding:"required,min=6"` +} + +// QuotaUpdateRequest 配额更新请求模型 +type QuotaUpdateRequest struct { + TotalTokens int64 `json:"total_tokens" binding:"required,min=0"` + TokenLimit int64 `json:"token_limit" binding:"required,min=1"` +} diff --git a/backend/services/ai_service.go b/backend/services/ai_service.go new file mode 100644 index 0000000..1481fe1 --- /dev/null +++ b/backend/services/ai_service.go @@ -0,0 +1,132 @@ +package services + +import ( + "bytes" + "encoding/json" + "fmt" + "io" + "net/http" + + "d:\项目\dobaochet\backend\config" +) + +// AIService AI服务结构体 +type AIService struct { + cfg *config.Config + client *http.Client +} + +// NewAIService 创建AI服务实例 +func NewAIService(cfg *config.Config) *AIService { + return &AIService{ + cfg: cfg, + client: &http.Client{}, + } +} + +// AIChatRequest AI聊天请求结构体 +type AIChatRequest struct { + Model string `json:"model"` + Messages []Message `json:"messages"` + MaxTokens int64 `json:"max_tokens"` + Temperature float64 `json:"temperature"` +} + +// Message 消息结构体 +type Message struct { + Role string `json:"role"` + Content string `json:"content"` +} + +// AIChatResponse AI聊天响应结构体 +type AIChatResponse struct { + ID string `json:"id"` + Object string `json:"object"` + Created int64 `json:"created"` + Model string `json:"model"` + Choices []Choice `json:"choices"` + Usage Usage `json:"usage"` +} + +// Choice 响应选项结构体 +type Choice struct { + Index int `json:"index"` + Message Message `json:"message"` + FinishReason string `json:"finish_reason"` +} + +// Usage 用量结构体 +type Usage struct { + PromptTokens int64 `json:"prompt_tokens"` + CompletionTokens int64 `json:"completion_tokens"` + TotalTokens int64 `json:"total_tokens"` +} + +// Chat 调用AI模型聊天 +func (s *AIService) Chat(messages []Message) (*AIChatResponse, error) { + // 构建请求 + reqBody := AIChatRequest{ + Model: s.cfg.AI.ModelName, + Messages: messages, + MaxTokens: s.cfg.AI.MaxTokens, + Temperature: s.cfg.AI.Temperature, + } + + // 序列化请求体 + reqBytes, err := json.Marshal(reqBody) + if err != nil { + return nil, fmt.Errorf("序列化请求失败: %w", err) + } + + // 创建HTTP请求 + req, err := http.NewRequest("POST", s.cfg.AI.APIURL, bytes.NewBuffer(reqBytes)) + if err != nil { + return nil, fmt.Errorf("创建请求失败: %w", err) + } + + // 设置请求头 + req.Header.Set("Content-Type", "application/json") + if s.cfg.AI.APIKey != "" { + req.Header.Set("Authorization", fmt.Sprintf("Bearer %s", s.cfg.AI.APIKey)) + } + + // 发送请求 + resp, err := s.client.Do(req) + if err != nil { + return nil, fmt.Errorf("发送请求失败: %w", err) + } + defer resp.Body.Close() + + // 读取响应 + respBody, err := io.ReadAll(resp.Body) + if err != nil { + return nil, fmt.Errorf("读取响应失败: %w", err) + } + + // 检查响应状态 + if resp.StatusCode != http.StatusOK { + return nil, fmt.Errorf("AI API返回错误: %s, 响应: %s", resp.Status, string(respBody)) + } + + // 解析响应 + var aiResp AIChatResponse + if err := json.Unmarshal(respBody, &aiResp); err != nil { + return nil, fmt.Errorf("解析响应失败: %w", err) + } + + return &aiResp, nil +} + +// CountTokens 估算令牌数量(简化实现,实际应使用更准确的方法) +func (s *AIService) CountTokens(text string) int64 { + // 简单估算:每个汉字算2个令牌,每个英文单词算1个令牌 + var count int64 + for _, r := range text { + if r > 127 { + count += 2 + } else { + count++ + } + } + return count / 2 // 平均估算 +} diff --git a/backend/services/quota_service.go b/backend/services/quota_service.go new file mode 100644 index 0000000..a348a57 --- /dev/null +++ b/backend/services/quota_service.go @@ -0,0 +1,162 @@ +package services + +import ( + "log" + "time" + + "d:\项目\dobaochet\backend\config" + "d:\项目\dobaochet\backend\db" + "d:\项目\dobaochet\backend\models" +) + +// QuotaService 配额管理服务 +type QuotaService struct { + cfg *config.Config +} + +// NewQuotaService 创建配额管理服务实例 +func NewQuotaService(cfg *config.Config) *QuotaService { + return &QuotaService{cfg: cfg} +} + +// StartQuotaResetJob 启动配额重置定时任务 +func (s *QuotaService) StartQuotaResetJob() { + // 立即执行一次重置检查 + s.ResetExpiredQuotas() + + // 每天检查一次 + ticker := time.NewTicker(24 * time.Hour) + go func() { + for range ticker.C { + s.ResetExpiredQuotas() + } + }() +} + +// ResetExpiredQuotas 重置过期的配额 +func (s *QuotaService) ResetExpiredQuotas() { + log.Println("开始检查并重置过期配额") + + // 查询所有需要重置配额的用户 + var users []models.User + result := db.GetDB().Preload("Quota").Where("quota.reset_at <= ?", time.Now()).Find(&users) + if result.Error != nil { + log.Printf("查询需要重置配额的用户失败: %v", result.Error) + return + } + + if len(users) == 0 { + log.Println("没有需要重置配额的用户") + return + } + + log.Printf("发现 %d 个用户需要重置配额", len(users)) + + // 批量重置配额 + tx := db.GetDB().Begin() + defer func() { + if r := recover(); r != nil { + tx.Rollback() + log.Printf("重置配额时发生 panic: %v", r) + } + }() + + for _, user := range users { + // 重置配额 + user.Quota.UsedTokens = 0 + user.Quota.ResetAt = time.Now().Add(s.cfg.Quota.ResetInterval) + + if result := tx.Save(&user.Quota); result.Error != nil { + tx.Rollback() + log.Printf("重置用户 %s 配额失败: %v", user.Username, result.Error) + return + } + + log.Printf("已重置用户 %s 的配额", user.Username) + } + + if err := tx.Commit().Error; err != nil { + log.Printf("提交配额重置事务失败: %v", err) + return + } + + log.Printf("成功重置 %d 个用户的配额", len(users)) +} + +// CheckAndUpdateQuota 检查并更新用户配额 +func (s *QuotaService) CheckAndUpdateQuota(userID uint, usedTokens int64) (bool, error) { + // 查询用户配额 + var quota models.Quota + result := db.GetDB().Where("user_id = ?", userID).First("a) + if result.Error != nil { + return false, result.Error + } + + // 检查配额是否过期 + if quota.ResetAt.Before(time.Now()) { + // 重置配额 + quota.UsedTokens = 0 + quota.ResetAt = time.Now().Add(s.cfg.Quota.ResetInterval) + } + + // 检查配额是否足够 + if quota.UsedTokens+usedTokens > quota.TotalTokens { + return false, nil + } + + // 更新配额 + quota.UsedTokens += usedTokens + if result := db.GetDB().Save("a); result.Error != nil { + return false, result.Error + } + + return true, nil +} + +// GetUserQuota 获取用户配额信息 +func (s *QuotaService) GetUserQuota(userID uint) (*models.Quota, error) { + var quota models.Quota + result := db.GetDB().Where("user_id = ?", userID).First("a) + if result.Error != nil { + return nil, result.Error + } + + // 检查并重置过期配额 + if quota.ResetAt.Before(time.Now()) { + quota.UsedTokens = 0 + quota.ResetAt = time.Now().Add(s.cfg.Quota.ResetInterval) + db.GetDB().Save("a) + } + + return "a, nil +} + +// UpdateUserQuota 更新用户配额 +func (s *QuotaService) UpdateUserQuota(userID uint, totalTokens, tokenLimit int64) error { + var quota models.Quota + result := db.GetDB().Where("user_id = ?", userID).First("a) + if result.Error != nil { + return result.Error + } + + // 更新配额 + quota.TotalTokens = totalTokens + quota.TokenLimit = tokenLimit + + return db.GetDB().Save("a).Error +} + +// ResetUserQuota 重置单个用户配额 +func (s *QuotaService) ResetUserQuota(userID uint) error { + var quota models.Quota + result := db.GetDB().Where("user_id = ?", userID).First("a) + if result.Error != nil { + return result.Error + } + + // 重置配额 + quota.UsedTokens = 0 + quota.ResetAt = time.Now().Add(s.cfg.Quota.ResetInterval) + + return db.GetDB().Save("a).Error +} diff --git a/backend/utils/auth.go b/backend/utils/auth.go new file mode 100644 index 0000000..f709767 --- /dev/null +++ b/backend/utils/auth.go @@ -0,0 +1,92 @@ +package utils + +import ( + "fmt" + "time" + + "github.com/golang-jwt/jwt/v5" + "golang.org/x/crypto/bcrypt" +) + +// JWTClaims JWT声明结构体 +type JWTClaims struct { + UserID uint `json:"user_id"` + Username string `json:"username"` + IsAdmin bool `json:"is_admin"` + jwt.RegisteredClaims +} + +// GenerateToken 生成JWT令牌 +func GenerateToken(userID uint, username string, isAdmin bool, secretKey string, expireTime time.Duration) (string, error) { + // 创建声明 + claims := JWTClaims{ + UserID: userID, + Username: username, + IsAdmin: isAdmin, + RegisteredClaims: jwt.RegisteredClaims{ + ExpiresAt: jwt.NewNumericDate(time.Now().Add(expireTime)), + IssuedAt: jwt.NewNumericDate(time.Now()), + NotBefore: jwt.NewNumericDate(time.Now()), + }, + } + + // 创建令牌 + token := jwt.NewWithClaims(jwt.SigningMethodHS256, claims) + + // 签名令牌 + tokenString, err := token.SignedString([]byte(secretKey)) + if err != nil { + return "", err + } + + return tokenString, nil +} + +// ParseToken 解析JWT令牌 +func ParseToken(tokenString string, secretKey string) (*JWTClaims, error) { + // 解析令牌 + token, err := jwt.ParseWithClaims(tokenString, &JWTClaims{}, func(token *jwt.Token) (interface{}, error) { + // 验证签名方法 + if _, ok := token.Method.(*jwt.SigningMethodHMAC); !ok { + return nil, fmt.Errorf("unexpected signing method: %v", token.Header["alg"]) + } + + return []byte(secretKey), nil + }) + + if err != nil { + return nil, err + } + + // 验证令牌并获取声明 + if claims, ok := token.Claims.(*JWTClaims); ok && token.Valid { + return claims, nil + } + + return nil, fmt.Errorf("invalid token") +} + +// HashPassword 密码加密 +func HashPassword(password string) (string, error) { + hash, err := bcrypt.GenerateFromPassword([]byte(password), bcrypt.DefaultCost) + if err != nil { + return "", err + } + + return string(hash), nil +} + +// VerifyPassword 密码验证 +func VerifyPassword(hashedPassword, password string) error { + return bcrypt.CompareHashAndPassword([]byte(hashedPassword), []byte(password)) +} + +// GenerateRandomString 生成随机字符串 +func GenerateRandomString(length int) string { + const charset = "abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789" + randomBytes := make([]byte, length) + for i := range randomBytes { + randomBytes[i] = charset[time.Now().UnixNano()%int64(len(charset))] + } + return string(randomBytes) +} diff --git a/env b/env deleted file mode 100644 index ac4e23b..0000000 --- a/env +++ /dev/null @@ -1,34 +0,0 @@ -# ========== 服务器基本配置 ========== -PORT=3001 -HOST=0.0.0.0 -SERVER_READ_TIMEOUT=60 -SERVER_WRITE_TIMEOUT=600 -SERVER_IDLE_TIMEOUT=120 -SERVER_GRACEFUL_SHUTDOWN_TIMEOUT=10 - -# ========== 节点/时区 ========== -IS_SLAVE=false -TZ=Asia/Shanghai - -# ========== 认证 ========== -AUTH_KEY=sk-123456 - -# ========== 数据库 ========== -# 与 compose 里的 mysql 服务一一对应 -DATABASE_DSN=root:123456@tcp(mysql:3306)/gpt-load?charset=utf8mb4&parseTime=True&loc=Local - -# ========== Redis ========== -# 与 compose 里的 redis 服务对应,使用 0 号库 -REDIS_DSN=redis://redis:6379/0 - -# ========== 并发 / CORS / 日志 ========== -MAX_CONCURRENT_REQUESTS=100 -ENABLE_CORS=true -ALLOWED_ORIGINS=* -ALLOWED_METHODS=GET,POST,PUT,DELETE,OPTIONS -ALLOWED_HEADERS=* -ALLOW_CREDENTIALS=false -LOG_LEVEL=info -LOG_FORMAT=text -LOG_ENABLE_FILE=true -LOG_FILE_PATH=./data/logs/app.log \ No newline at end of file diff --git a/index.html b/index.html new file mode 100644 index 0000000..0104aa0 --- /dev/null +++ b/index.html @@ -0,0 +1,79 @@ + + +
+ + +