knowledge_base_controller.go 5.6 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200
  1. package controller
  2. import (
  3. "log"
  4. "net/http"
  5. "strconv"
  6. "github.com/2930134478/AI-CS/backend/service"
  7. "github.com/gin-gonic/gin"
  8. )
  9. // KnowledgeBaseController 知识库控制器
  10. type KnowledgeBaseController struct {
  11. knowledgeBaseService *service.KnowledgeBaseService
  12. embeddingConfigService *service.EmbeddingConfigService
  13. }
  14. // NewKnowledgeBaseController 创建知识库控制器实例
  15. func NewKnowledgeBaseController(knowledgeBaseService *service.KnowledgeBaseService, embeddingConfigService *service.EmbeddingConfigService) *KnowledgeBaseController {
  16. return &KnowledgeBaseController{
  17. knowledgeBaseService: knowledgeBaseService,
  18. embeddingConfigService: embeddingConfigService,
  19. }
  20. }
  21. // checkKBAccess 校验当前用户是否允许使用知识库(请求头须带 X-User-Id;未带则放行以兼容旧前端)
  22. func (c *KnowledgeBaseController) checkKBAccess(ctx *gin.Context) bool {
  23. userID := getUserIDFromHeader(ctx)
  24. if userID == 0 {
  25. return true
  26. }
  27. if err := c.embeddingConfigService.CheckKnowledgeBaseAccess(userID); err != nil {
  28. ctx.JSON(http.StatusForbidden, gin.H{"error": err.Error()})
  29. return false
  30. }
  31. return true
  32. }
  33. // ListKnowledgeBases 获取知识库列表
  34. func (c *KnowledgeBaseController) ListKnowledgeBases(ctx *gin.Context) {
  35. if !c.checkKBAccess(ctx) {
  36. return
  37. }
  38. kbs, err := c.knowledgeBaseService.ListKnowledgeBases()
  39. if err != nil {
  40. log.Printf("获取知识库列表失败: %v", err)
  41. ctx.JSON(http.StatusInternalServerError, gin.H{"error": "获取知识库列表失败"})
  42. return
  43. }
  44. ctx.JSON(http.StatusOK, gin.H{
  45. "knowledge_bases": kbs,
  46. })
  47. }
  48. // GetKnowledgeBase 获取知识库详情
  49. func (c *KnowledgeBaseController) GetKnowledgeBase(ctx *gin.Context) {
  50. if !c.checkKBAccess(ctx) {
  51. return
  52. }
  53. idStr := ctx.Param("id")
  54. id, err := strconv.ParseUint(idStr, 10, 64)
  55. if err != nil || id == 0 {
  56. ctx.JSON(http.StatusBadRequest, gin.H{"error": "知识库 ID 不合法"})
  57. return
  58. }
  59. kb, err := c.knowledgeBaseService.GetKnowledgeBase(uint(id))
  60. if err != nil {
  61. log.Printf("获取知识库详情失败: %v", err)
  62. ctx.JSON(http.StatusNotFound, gin.H{"error": err.Error()})
  63. return
  64. }
  65. ctx.JSON(http.StatusOK, kb)
  66. }
  67. // CreateKnowledgeBase 创建知识库
  68. func (c *KnowledgeBaseController) CreateKnowledgeBase(ctx *gin.Context) {
  69. if !c.checkKBAccess(ctx) {
  70. return
  71. }
  72. var req struct {
  73. Name string `json:"name" binding:"required"`
  74. Description string `json:"description"`
  75. }
  76. if err := ctx.ShouldBindJSON(&req); err != nil {
  77. ctx.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
  78. return
  79. }
  80. kb, err := c.knowledgeBaseService.CreateKnowledgeBase(service.CreateKnowledgeBaseInput{
  81. Name: req.Name,
  82. Description: req.Description,
  83. })
  84. if err != nil {
  85. log.Printf("创建知识库失败: %v", err)
  86. ctx.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
  87. return
  88. }
  89. ctx.JSON(http.StatusOK, kb)
  90. }
  91. // UpdateKnowledgeBase 更新知识库
  92. func (c *KnowledgeBaseController) UpdateKnowledgeBase(ctx *gin.Context) {
  93. if !c.checkKBAccess(ctx) {
  94. return
  95. }
  96. idStr := ctx.Param("id")
  97. id, err := strconv.ParseUint(idStr, 10, 64)
  98. if err != nil || id == 0 {
  99. ctx.JSON(http.StatusBadRequest, gin.H{"error": "知识库 ID 不合法"})
  100. return
  101. }
  102. var req struct {
  103. Name *string `json:"name"`
  104. Description *string `json:"description"`
  105. RAGEnabled *bool `json:"rag_enabled"`
  106. }
  107. if err := ctx.ShouldBindJSON(&req); err != nil {
  108. ctx.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
  109. return
  110. }
  111. kb, err := c.knowledgeBaseService.UpdateKnowledgeBase(uint(id), service.UpdateKnowledgeBaseInput{
  112. Name: req.Name,
  113. Description: req.Description,
  114. RAGEnabled: req.RAGEnabled,
  115. })
  116. if err != nil {
  117. log.Printf("更新知识库失败: %v", err)
  118. ctx.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
  119. return
  120. }
  121. ctx.JSON(http.StatusOK, kb)
  122. }
  123. // DeleteKnowledgeBase 删除知识库
  124. func (c *KnowledgeBaseController) DeleteKnowledgeBase(ctx *gin.Context) {
  125. if !c.checkKBAccess(ctx) {
  126. return
  127. }
  128. idStr := ctx.Param("id")
  129. id, err := strconv.ParseUint(idStr, 10, 64)
  130. if err != nil || id == 0 {
  131. ctx.JSON(http.StatusBadRequest, gin.H{"error": "知识库 ID 不合法"})
  132. return
  133. }
  134. if err := c.knowledgeBaseService.DeleteKnowledgeBase(uint(id)); err != nil {
  135. log.Printf("删除知识库失败: %v", err)
  136. ctx.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
  137. return
  138. }
  139. ctx.JSON(http.StatusOK, gin.H{"message": "删除成功"})
  140. }
  141. // UpdateKnowledgeBaseRAGEnabled 仅更新知识库「参与 RAG」开关。
  142. func (c *KnowledgeBaseController) UpdateKnowledgeBaseRAGEnabled(ctx *gin.Context) {
  143. if !c.checkKBAccess(ctx) {
  144. return
  145. }
  146. idStr := ctx.Param("id")
  147. id, err := strconv.ParseUint(idStr, 10, 64)
  148. if err != nil || id == 0 {
  149. ctx.JSON(http.StatusBadRequest, gin.H{"error": "知识库 ID 不合法"})
  150. return
  151. }
  152. var req struct {
  153. RAGEnabled bool `json:"rag_enabled"`
  154. }
  155. if err := ctx.ShouldBindJSON(&req); err != nil {
  156. ctx.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
  157. return
  158. }
  159. kb, err := c.knowledgeBaseService.UpdateKnowledgeBase(uint(id), service.UpdateKnowledgeBaseInput{
  160. RAGEnabled: &req.RAGEnabled,
  161. })
  162. if err != nil {
  163. log.Printf("更新知识库 RAG 开关失败: %v", err)
  164. ctx.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
  165. return
  166. }
  167. ctx.JSON(http.StatusOK, kb)
  168. }
  169. // ListDocumentsByKnowledgeBase 获取知识库的文档列表
  170. func (c *KnowledgeBaseController) ListDocumentsByKnowledgeBase(ctx *gin.Context) {
  171. if !c.checkKBAccess(ctx) {
  172. return
  173. }
  174. // 这个功能由 DocumentController 实现,这里可以重定向或调用
  175. ctx.JSON(http.StatusOK, gin.H{"message": "请使用 /documents?knowledge_base_id=:id"})
  176. }