Vous ne pouvez pas sélectionner plus de 25 sujets Les noms de sujets doivent commencer par une lettre ou un nombre, peuvent contenir des tirets ('-') et peuvent comporter jusqu'à 35 caractères.
 
 
 

158 lignes
5.4 KiB

  1. package controller
  2. import (
  3. "context"
  4. "fmt"
  5. "io"
  6. "net/http"
  7. "strings"
  8. "github.com/QuantumNous/new-api/common"
  9. "github.com/QuantumNous/new-api/constant"
  10. "github.com/QuantumNous/new-api/logger"
  11. "github.com/QuantumNous/new-api/model"
  12. "github.com/QuantumNous/new-api/service"
  13. "github.com/gin-gonic/gin"
  14. )
  15. func assetProxyError(c *gin.Context, status int, errType, message string) {
  16. c.JSON(status, gin.H{
  17. "error": gin.H{
  18. "message": message,
  19. "type": errType,
  20. },
  21. })
  22. }
  23. var doubaoAssetAutoGroups = service.GetUserAutoGroup
  24. func effectiveDoubaoAssetGroup(c *gin.Context) string {
  25. if group := strings.TrimSpace(common.GetContextKeyString(c, constant.ContextKeyUsingGroup)); group != "" {
  26. return group
  27. }
  28. return strings.TrimSpace(common.GetContextKeyString(c, constant.ContextKeyTokenGroup))
  29. }
  30. func concreteDoubaoAssetGroupsForRequest(c *gin.Context, autoGroups func(string) []string) []string {
  31. group := effectiveDoubaoAssetGroup(c)
  32. if group != "auto" {
  33. if group == "" {
  34. return nil
  35. }
  36. return []string{group}
  37. }
  38. userGroup := common.GetContextKeyString(c, constant.ContextKeyUserGroup)
  39. groups := autoGroups(userGroup)
  40. concreteGroups := make([]string, 0, len(groups))
  41. for _, candidate := range groups {
  42. candidate = strings.TrimSpace(candidate)
  43. if candidate == "" || candidate == "auto" {
  44. continue
  45. }
  46. concreteGroups = append(concreteGroups, candidate)
  47. }
  48. return concreteGroups
  49. }
  50. func DoubaoAssetProxy(c *gin.Context) {
  51. actionRaw := strings.TrimSpace(c.Query("Action"))
  52. if actionRaw == "" {
  53. assetProxyError(c, http.StatusBadRequest, service.AssetErrorInvalidRequest, "Action query parameter is required")
  54. return
  55. }
  56. action, ok := service.ParseAssetAction(actionRaw)
  57. if !ok {
  58. assetProxyError(c, http.StatusBadRequest, service.AssetErrorInvalidRequest, fmt.Sprintf("unsupported asset Action: %s", actionRaw))
  59. return
  60. }
  61. version := strings.TrimSpace(c.Query("Version"))
  62. if version == "" {
  63. version = "2024-01-01"
  64. }
  65. rawBody, err := io.ReadAll(c.Request.Body)
  66. if err != nil {
  67. assetProxyError(c, http.StatusBadRequest, service.AssetErrorInvalidRequest, fmt.Sprintf("failed to read request body: %v", err))
  68. return
  69. }
  70. body := map[string]any{}
  71. if len(strings.TrimSpace(string(rawBody))) > 0 {
  72. if err := common.Unmarshal(rawBody, &body); err != nil {
  73. assetProxyError(c, http.StatusBadRequest, service.AssetErrorInvalidRequest, fmt.Sprintf("invalid JSON body: %v", err))
  74. return
  75. }
  76. }
  77. userID := c.GetInt("id")
  78. var lastAssetErr *service.AssetError
  79. for _, group := range concreteDoubaoAssetGroupsForRequest(c, doubaoAssetAutoGroups) {
  80. channel, adapter, assetErr := service.ResolveAssetChannelForOperation(userID, group, action.Operation)
  81. if assetErr != nil {
  82. lastAssetErr = assetErr
  83. logger.LogError(c.Request.Context(), fmt.Sprintf("Failed to resolve asset channel for group %s action %s: %s", group, action.Action, assetErr.Message))
  84. if assetErr.Type == service.AssetErrorChannelNotFound {
  85. continue
  86. }
  87. break
  88. }
  89. if channel == nil || adapter == nil {
  90. continue
  91. }
  92. assetRequest := service.AssetRequest{
  93. Action: action,
  94. Version: version,
  95. Body: body,
  96. RawBody: rawBody,
  97. }
  98. if service.IsChinaMobileAssetChannel(channel) {
  99. if service.IsAssetGroupOperation(action.Operation) {
  100. assetProxyError(c, http.StatusForbidden, service.AssetErrorOperationNotSupported, "China Mobile asset group APIs are managed by the platform")
  101. return
  102. }
  103. groupID, assetErr := service.GetOrCreateChinaMobileUserAssetGroup(c.Request.Context(), userID, channel, func(ctx context.Context, userID int, channel *model.Channel) (string, *service.AssetError) {
  104. return service.CreateChinaMobileUserAssetGroup(ctx, userID, adapter, channel)
  105. })
  106. if assetErr != nil {
  107. assetProxyError(c, assetErr.HTTPStatus, assetErr.Type, assetErr.Message)
  108. return
  109. }
  110. service.ScopeChinaMobileAssetRequest(&assetRequest, groupID)
  111. if action.Operation == service.AssetOperationAssetGet || action.Operation == service.AssetOperationAssetUpdate || action.Operation == service.AssetOperationAssetDelete {
  112. if assetErr := service.RequireChinaMobileAssetOwnership(c.Request.Context(), adapter, channel, assetRequest, groupID); assetErr != nil {
  113. assetProxyError(c, assetErr.HTTPStatus, assetErr.Type, assetErr.Message)
  114. return
  115. }
  116. }
  117. }
  118. resp, assetErr := adapter.DoAssetRequest(c.Request.Context(), channel, assetRequest)
  119. if assetErr != nil {
  120. logger.LogError(c.Request.Context(), fmt.Sprintf("Failed to proxy asset action %s via channel %d type %d: %s", action.Action, channel.Id, channel.Type, assetErr.Message))
  121. assetProxyError(c, assetErr.HTTPStatus, assetErr.Type, assetErr.Message)
  122. return
  123. }
  124. copyDoubaoAssetResponseHeaders(c, resp.Header)
  125. c.Data(resp.StatusCode, "application/json", resp.Body)
  126. return
  127. }
  128. if lastAssetErr != nil {
  129. assetProxyError(c, lastAssetErr.HTTPStatus, lastAssetErr.Type, lastAssetErr.Message)
  130. return
  131. }
  132. assetProxyError(c, http.StatusBadGateway, service.AssetErrorChannelNotFound, "no available asset channel supports requested operation")
  133. }
  134. func copyDoubaoAssetResponseHeaders(c *gin.Context, headers http.Header) {
  135. for key, values := range headers {
  136. switch http.CanonicalHeaderKey(key) {
  137. case "Connection", "Content-Length", "Keep-Alive", "Proxy-Authenticate", "Proxy-Authorization", "Te", "Trailer", "Transfer-Encoding", "Upgrade":
  138. continue
  139. }
  140. for _, value := range values {
  141. c.Writer.Header().Add(key, value)
  142. }
  143. }
  144. }