import_controller.go 4.5 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163
  1. package controller
  2. import (
  3. "context"
  4. "log"
  5. "net/http"
  6. "os"
  7. "path/filepath"
  8. "strconv"
  9. "strings"
  10. "github.com/2930134478/AI-CS/backend/service"
  11. "github.com/gin-gonic/gin"
  12. )
  13. // ImportController 导入控制器
  14. type ImportController struct {
  15. importService *service.ImportService
  16. embeddingConfigService *service.EmbeddingConfigService
  17. }
  18. // NewImportController 创建导入控制器实例
  19. func NewImportController(importService *service.ImportService, embeddingConfigService *service.EmbeddingConfigService) *ImportController {
  20. return &ImportController{
  21. importService: importService,
  22. embeddingConfigService: embeddingConfigService,
  23. }
  24. }
  25. func (c *ImportController) checkKBAccess(ctx *gin.Context) bool {
  26. userID := getUserIDFromHeader(ctx)
  27. if userID == 0 {
  28. // ⚠️ 修复:改为拒绝访问,而不是允许
  29. ctx.JSON(http.StatusUnauthorized, gin.H{"error": "未授权访问,请提供 X-User-Id 请求头"})
  30. return false
  31. }
  32. if err := c.embeddingConfigService.CheckKnowledgeBaseAccess(userID); err != nil {
  33. ctx.JSON(http.StatusForbidden, gin.H{"error": err.Error()})
  34. return false
  35. }
  36. return true
  37. }
  38. // ImportDocuments 批量导入文档(文件上传)
  39. func (c *ImportController) ImportDocuments(ctx *gin.Context) {
  40. if !c.checkKBAccess(ctx) {
  41. return
  42. }
  43. // 获取知识库 ID
  44. kbIDStr := ctx.PostForm("knowledge_base_id")
  45. if kbIDStr == "" {
  46. ctx.JSON(http.StatusBadRequest, gin.H{"error": "知识库 ID 不能为空"})
  47. return
  48. }
  49. kbID, err := strconv.ParseUint(kbIDStr, 10, 64)
  50. if err != nil || kbID == 0 {
  51. ctx.JSON(http.StatusBadRequest, gin.H{"error": "知识库 ID 不合法"})
  52. return
  53. }
  54. // 获取上传的文件
  55. form, err := ctx.MultipartForm()
  56. if err != nil {
  57. ctx.JSON(http.StatusBadRequest, gin.H{"error": "获取文件失败"})
  58. return
  59. }
  60. files := form.File["files"]
  61. if len(files) == 0 {
  62. ctx.JSON(http.StatusBadRequest, gin.H{"error": "未上传文件"})
  63. return
  64. }
  65. // ⚠️ 添加:文件类型验证
  66. allowedExts := map[string]bool{
  67. ".md": true,
  68. ".txt": true,
  69. ".pdf": true,
  70. ".doc": true,
  71. ".docx": true,
  72. }
  73. // 保存文件到临时目录
  74. filePaths := make([]string, 0, len(files))
  75. for _, file := range files {
  76. // ⚠️ 添加:验证文件类型
  77. ext := strings.ToLower(filepath.Ext(file.Filename))
  78. if !allowedExts[ext] {
  79. log.Printf("不支持的文件类型: %s (扩展名: %s)", file.Filename, ext)
  80. continue
  81. }
  82. // ⚠️ 添加:清理文件名,防止路径遍历攻击
  83. safeFilename := filepath.Base(file.Filename)
  84. safeFilename = strings.ReplaceAll(safeFilename, "..", "")
  85. safeFilename = strings.ReplaceAll(safeFilename, "/", "")
  86. safeFilename = strings.ReplaceAll(safeFilename, "\\", "")
  87. // 限制文件名长度
  88. if len(safeFilename) > 255 {
  89. safeFilename = safeFilename[:255]
  90. }
  91. // 保存文件
  92. filePath := "/tmp/" + safeFilename
  93. if err := ctx.SaveUploadedFile(file, filePath); err != nil {
  94. log.Printf("保存文件失败: %v", err)
  95. continue
  96. }
  97. filePaths = append(filePaths, filePath)
  98. }
  99. if len(filePaths) == 0 {
  100. ctx.JSON(http.StatusBadRequest, gin.H{"error": "没有有效的文件(所有文件都被拒绝或保存失败)"})
  101. return
  102. }
  103. // ⚠️ 添加:导入后清理临时文件
  104. defer func() {
  105. for _, path := range filePaths {
  106. if err := os.Remove(path); err != nil {
  107. log.Printf("清理临时文件失败: %v", err)
  108. }
  109. }
  110. }()
  111. // 导入文件
  112. result, err := c.importService.ImportFiles(context.Background(), uint(kbID), filePaths)
  113. if err != nil {
  114. log.Printf("导入文件失败: %v", err)
  115. ctx.JSON(http.StatusInternalServerError, gin.H{"error": "批量导入失败: " + err.Error()})
  116. return
  117. }
  118. result.Message = "导入完成"
  119. ctx.JSON(http.StatusOK, result)
  120. }
  121. // ImportFromURLs 批量导入文档(URL 爬取)
  122. func (c *ImportController) ImportFromURLs(ctx *gin.Context) {
  123. if !c.checkKBAccess(ctx) {
  124. return
  125. }
  126. var req struct {
  127. KnowledgeBaseID uint `json:"knowledge_base_id" binding:"required"`
  128. URLs []string `json:"urls" binding:"required"`
  129. }
  130. if err := ctx.ShouldBindJSON(&req); err != nil {
  131. ctx.JSON(http.StatusBadRequest, gin.H{"error": "请求参数错误: " + err.Error()})
  132. return
  133. }
  134. result, err := c.importService.ImportFromUrls(context.Background(), req.KnowledgeBaseID, req.URLs)
  135. if err != nil {
  136. log.Printf("导入 URL 失败: %v", err)
  137. ctx.JSON(http.StatusInternalServerError, gin.H{"error": "批量导入失败: " + err.Error()})
  138. return
  139. }
  140. result.Message = "导入完成"
  141. ctx.JSON(http.StatusOK, result)
  142. }