knowledge_base_controller.go 6.2 KB

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