diff --git a/controller/swag_video.go b/controller/swag_video.go deleted file mode 100644 index 68dd6345..00000000 --- a/controller/swag_video.go +++ /dev/null @@ -1,136 +0,0 @@ -package controller - -import ( - "github.com/gin-gonic/gin" -) - -// VideoGenerations -// @Summary 生成视频 -// @Description 调用视频生成接口生成视频 -// @Description 支持多种视频生成服务: -// @Description - 可灵AI (Kling): https://app.klingai.com/cn/dev/document-api/apiReference/commonInfo -// @Description - 即梦 (Jimeng): https://www.volcengine.com/docs/85621/1538636 -// @Tags Video -// @Accept json -// @Produce json -// @Param Authorization header string true "用户认证令牌 (Aeess-Token: sk-xxxx)" -// @Param request body dto.VideoRequest true "视频生成请求参数" -// @Failure 400 {object} dto.OpenAIError "请求参数错误" -// @Failure 401 {object} dto.OpenAIError "未授权" -// @Failure 403 {object} dto.OpenAIError "无权限" -// @Failure 500 {object} dto.OpenAIError "服务器内部错误" -// @Router /v1/video/generations [post] -func VideoGenerations(c *gin.Context) { -} - -// VideoGenerationsTaskId -// @Summary 查询视频 -// @Description 根据任务ID查询视频生成任务的状态和结果 -// @Tags Video -// @Accept json -// @Produce json -// @Security BearerAuth -// @Param task_id path string true "Task ID" -// @Success 200 {object} dto.VideoTaskResponse "任务状态和结果" -// @Failure 400 {object} dto.OpenAIError "请求参数错误" -// @Failure 401 {object} dto.OpenAIError "未授权" -// @Failure 403 {object} dto.OpenAIError "无权限" -// @Failure 500 {object} dto.OpenAIError "服务器内部错误" -// @Router /v1/video/generations/{task_id} [get] -func VideoGenerationsTaskId(c *gin.Context) { -} - -// KlingText2VideoGenerations -// @Summary 可灵文生视频 -// @Description 调用可灵AI文生视频接口,生成视频内容 -// @Tags Video -// @Accept json -// @Produce json -// @Param Authorization header string true "用户认证令牌 (Aeess-Token: sk-xxxx)" -// @Param request body KlingText2VideoRequest true "视频生成请求参数" -// @Success 200 {object} dto.VideoTaskResponse "任务状态和结果" -// @Failure 400 {object} dto.OpenAIError "请求参数错误" -// @Failure 401 {object} dto.OpenAIError "未授权" -// @Failure 403 {object} dto.OpenAIError "无权限" -// @Failure 500 {object} dto.OpenAIError "服务器内部错误" -// @Router /kling/v1/videos/text2video [post] -func KlingText2VideoGenerations(c *gin.Context) { -} - -type KlingText2VideoRequest struct { - ModelName string `json:"model_name,omitempty" example:"kling-v1"` - Prompt string `json:"prompt" binding:"required" example:"A cat playing piano in the garden"` - NegativePrompt string `json:"negative_prompt,omitempty" example:"blurry, low quality"` - CfgScale float64 `json:"cfg_scale,omitempty" example:"0.7"` - Mode string `json:"mode,omitempty" example:"std"` - CameraControl *KlingCameraControl `json:"camera_control,omitempty"` - AspectRatio string `json:"aspect_ratio,omitempty" example:"16:9"` - Duration string `json:"duration,omitempty" example:"5"` - CallbackURL string `json:"callback_url,omitempty" example:"https://your.domain/callback"` - ExternalTaskId string `json:"external_task_id,omitempty" example:"custom-task-001"` -} - -type KlingCameraControl struct { - Type string `json:"type,omitempty" example:"simple"` - Config *KlingCameraConfig `json:"config,omitempty"` -} - -type KlingCameraConfig struct { - Horizontal float64 `json:"horizontal,omitempty" example:"2.5"` - Vertical float64 `json:"vertical,omitempty" example:"0"` - Pan float64 `json:"pan,omitempty" example:"0"` - Tilt float64 `json:"tilt,omitempty" example:"0"` - Roll float64 `json:"roll,omitempty" example:"0"` - Zoom float64 `json:"zoom,omitempty" example:"0"` -} - -// KlingImage2VideoGenerations -// @Summary 可灵官方-图生视频 -// @Description 调用可灵AI图生视频接口,生成视频内容 -// @Tags Video -// @Accept json -// @Produce json -// @Param Authorization header string true "用户认证令牌 (Aeess-Token: sk-xxxx)" -// @Param request body KlingImage2VideoRequest true "图生视频请求参数" -// @Success 200 {object} dto.VideoTaskResponse "任务状态和结果" -// @Failure 400 {object} dto.OpenAIError "请求参数错误" -// @Failure 401 {object} dto.OpenAIError "未授权" -// @Failure 403 {object} dto.OpenAIError "无权限" -// @Failure 500 {object} dto.OpenAIError "服务器内部错误" -// @Router /kling/v1/videos/image2video [post] -func KlingImage2VideoGenerations(c *gin.Context) { -} - -type KlingImage2VideoRequest struct { - ModelName string `json:"model_name,omitempty" example:"kling-v2-master"` - Image string `json:"image" binding:"required" example:"https://h2.inkwai.com/bs2/upload-ylab-stunt/se/ai_portal_queue_mmu_image_upscale_aiweb/3214b798-e1b4-4b00-b7af-72b5b0417420_raw_image_0.jpg"` - Prompt string `json:"prompt,omitempty" example:"A cat playing piano in the garden"` - NegativePrompt string `json:"negative_prompt,omitempty" example:"blurry, low quality"` - CfgScale float64 `json:"cfg_scale,omitempty" example:"0.7"` - Mode string `json:"mode,omitempty" example:"std"` - CameraControl *KlingCameraControl `json:"camera_control,omitempty"` - AspectRatio string `json:"aspect_ratio,omitempty" example:"16:9"` - Duration string `json:"duration,omitempty" example:"5"` - CallbackURL string `json:"callback_url,omitempty" example:"https://your.domain/callback"` - ExternalTaskId string `json:"external_task_id,omitempty" example:"custom-task-002"` -} - -// KlingImage2videoTaskId godoc -// @Summary 可灵任务查询--图生视频 -// @Description Query the status and result of a Kling video generation task by task ID -// @Tags Origin -// @Accept json -// @Produce json -// @Param task_id path string true "Task ID" -// @Router /kling/v1/videos/image2video/{task_id} [get] -func KlingImage2videoTaskId(c *gin.Context) {} - -// KlingText2videoTaskId godoc -// @Summary 可灵任务查询--文生视频 -// @Description Query the status and result of a Kling text-to-video generation task by task ID -// @Tags Origin -// @Accept json -// @Produce json -// @Param task_id path string true "Task ID" -// @Router /kling/v1/videos/text2video/{task_id} [get] -func KlingText2videoTaskId(c *gin.Context) {} diff --git a/controller/task_video.go b/controller/task_video.go deleted file mode 100644 index 0c9f5e8d..00000000 --- a/controller/task_video.go +++ /dev/null @@ -1,327 +0,0 @@ -package controller - -import ( - "context" - "encoding/json" - "fmt" - "io" - "time" - - "github.com/QuantumNous/new-api/common" - "github.com/QuantumNous/new-api/constant" - "github.com/QuantumNous/new-api/dto" - "github.com/QuantumNous/new-api/logger" - "github.com/QuantumNous/new-api/model" - "github.com/QuantumNous/new-api/relay" - "github.com/QuantumNous/new-api/relay/channel" - relaycommon "github.com/QuantumNous/new-api/relay/common" - "github.com/QuantumNous/new-api/setting/ratio_setting" -) - -func UpdateVideoTaskAll(ctx context.Context, platform constant.TaskPlatform, taskChannelM map[int][]string, taskM map[string]*model.Task) error { - for channelId, taskIds := range taskChannelM { - if err := updateVideoTaskAll(ctx, platform, channelId, taskIds, taskM); err != nil { - logger.LogError(ctx, fmt.Sprintf("Channel #%d failed to update video async tasks: %s", channelId, err.Error())) - } - } - return nil -} - -func updateVideoTaskAll(ctx context.Context, platform constant.TaskPlatform, channelId int, taskIds []string, taskM map[string]*model.Task) error { - logger.LogInfo(ctx, fmt.Sprintf("Channel #%d pending video tasks: %d", channelId, len(taskIds))) - if len(taskIds) == 0 { - return nil - } - cacheGetChannel, err := model.CacheGetChannel(channelId) - if err != nil { - errUpdate := model.TaskBulkUpdate(taskIds, map[string]any{ - "fail_reason": fmt.Sprintf("Failed to get channel info, channel ID: %d", channelId), - "status": "FAILURE", - "progress": "100%", - }) - if errUpdate != nil { - common.SysLog(fmt.Sprintf("UpdateVideoTask error: %v", errUpdate)) - } - return fmt.Errorf("CacheGetChannel failed: %w", err) - } - adaptor := relay.GetTaskAdaptor(platform) - if adaptor == nil { - return fmt.Errorf("video adaptor not found") - } - info := &relaycommon.RelayInfo{} - info.ChannelMeta = &relaycommon.ChannelMeta{ - ChannelBaseUrl: cacheGetChannel.GetBaseURL(), - } - info.ApiKey = cacheGetChannel.Key - adaptor.Init(info) - for _, taskId := range taskIds { - if err := updateVideoSingleTask(ctx, adaptor, cacheGetChannel, taskId, taskM); err != nil { - logger.LogError(ctx, fmt.Sprintf("Failed to update video task %s: %s", taskId, err.Error())) - } - } - return nil -} - -func updateVideoSingleTask(ctx context.Context, adaptor channel.TaskAdaptor, channel *model.Channel, taskId string, taskM map[string]*model.Task) error { - baseURL := constant.ChannelBaseURLs[channel.Type] - if channel.GetBaseURL() != "" { - baseURL = channel.GetBaseURL() - } - proxy := channel.GetSetting().Proxy - - task := taskM[taskId] - if task == nil { - logger.LogError(ctx, fmt.Sprintf("Task %s not found in taskM", taskId)) - return fmt.Errorf("task %s not found", taskId) - } - key := channel.Key - - privateData := task.PrivateData - if privateData.Key != "" { - key = privateData.Key - } - resp, err := adaptor.FetchTask(baseURL, key, map[string]any{ - "task_id": taskId, - "action": task.Action, - }, proxy) - if err != nil { - return fmt.Errorf("fetchTask failed for task %s: %w", taskId, err) - } - //if resp.StatusCode != http.StatusOK { - //return fmt.Errorf("get Video Task status code: %d", resp.StatusCode) - //} - defer resp.Body.Close() - responseBody, err := io.ReadAll(resp.Body) - if err != nil { - return fmt.Errorf("readAll failed for task %s: %w", taskId, err) - } - - logger.LogDebug(ctx, "UpdateVideoSingleTask response: %s", responseBody) - - taskResult := &relaycommon.TaskInfo{} - // try parse as New API response format - var responseItems dto.TaskResponse[model.Task] - if err = common.Unmarshal(responseBody, &responseItems); err == nil && responseItems.IsSuccess() { - logger.LogDebug(ctx, "UpdateVideoSingleTask parsed as new api response format: %+v", responseItems) - t := responseItems.Data - taskResult.TaskID = t.TaskID - taskResult.Status = string(t.Status) - taskResult.Url = t.FailReason - taskResult.Progress = t.Progress - taskResult.Reason = t.FailReason - task.Data = t.Data - } else if taskResult, err = adaptor.ParseTaskResult(responseBody); err != nil { - return fmt.Errorf("parseTaskResult failed for task %s: %w", taskId, err) - } else { - task.Data = redactVideoResponseBody(responseBody) - } - - logger.LogDebug(ctx, "UpdateVideoSingleTask taskResult: %+v", taskResult) - - now := time.Now().Unix() - if taskResult.Status == "" { - //return fmt.Errorf("task %s status is empty", taskId) - taskResult = relaycommon.FailTaskInfo("upstream returned empty status") - } - - // 记录原本的状态,防止重复退款 - shouldRefund := false - quota := task.Quota - preStatus := task.Status - - task.Status = model.TaskStatus(taskResult.Status) - switch taskResult.Status { - case model.TaskStatusSubmitted: - task.Progress = "10%" - case model.TaskStatusQueued: - task.Progress = "20%" - case model.TaskStatusInProgress: - task.Progress = "30%" - if task.StartTime == 0 { - task.StartTime = now - } - case model.TaskStatusSuccess: - task.Progress = "100%" - if task.FinishTime == 0 { - task.FinishTime = now - } - if !(len(taskResult.Url) > 5 && taskResult.Url[:5] == "data:") { - task.FailReason = taskResult.Url - } - - // 如果返回了 total_tokens 并且配置了模型倍率(非固定价格),则重新计费 - if taskResult.TotalTokens > 0 { - // 获取模型名称 - var taskData map[string]interface{} - if err := json.Unmarshal(task.Data, &taskData); err == nil { - if modelName, ok := taskData["model"].(string); ok && modelName != "" { - // 获取模型价格和倍率 - modelRatio, hasRatioSetting, _ := ratio_setting.GetModelRatio(modelName) - // 只有配置了倍率(非固定价格)时才按 token 重新计费 - if hasRatioSetting && modelRatio > 0 { - // 获取用户和组的倍率信息 - group := task.Group - if group == "" { - user, err := model.GetUserById(task.UserId, false) - if err == nil { - group = user.Group - } - } - if group != "" { - groupRatio := ratio_setting.GetGroupRatio(group) - userGroupRatio, hasUserGroupRatio := ratio_setting.GetGroupGroupRatio(group, group) - - var finalGroupRatio float64 - if hasUserGroupRatio { - finalGroupRatio = userGroupRatio - } else { - finalGroupRatio = groupRatio - } - - // 计算实际应扣费额度: totalTokens * modelRatio * groupRatio(饱和转换,防止溢出成负数) - actualQuota, clamp := common.QuotaFromFloatChecked(float64(taskResult.TotalTokens) * modelRatio * finalGroupRatio) - if clamp != nil { - logger.LogWarn(ctx, fmt.Sprintf("quota saturation on video task %s: op=%s kind=%s original=%g clamped=%d user=%d", - task.TaskID, clamp.Op, clamp.Kind, clamp.Original, clamp.Clamped, task.UserId)) - } - - // 计算差额 - preConsumedQuota := task.Quota - quotaDelta := actualQuota - preConsumedQuota - - if quotaDelta > 0 { - // 需要补扣费 - logger.LogInfo(ctx, fmt.Sprintf("视频任务 %s 预扣费后补扣费:%s(实际消耗:%s,预扣费:%s,tokens:%d)", - task.TaskID, - logger.LogQuota(quotaDelta), - logger.LogQuota(actualQuota), - logger.LogQuota(preConsumedQuota), - taskResult.TotalTokens, - )) - if err := model.DecreaseUserQuota(task.UserId, quotaDelta, false); err != nil { - logger.LogError(ctx, fmt.Sprintf("补扣费失败: %s", err.Error())) - } else { - model.UpdateUserUsedQuotaAndRequestCount(task.UserId, quotaDelta) - model.UpdateChannelUsedQuota(task.ChannelId, quotaDelta) - task.Quota = actualQuota // 更新任务记录的实际扣费额度 - - // 记录消费日志 - logContent := fmt.Sprintf("视频任务成功补扣费,模型倍率 %.2f,分组倍率 %.2f,tokens %d,预扣费 %s,实际扣费 %s,补扣费 %s", - modelRatio, finalGroupRatio, taskResult.TotalTokens, - logger.LogQuota(preConsumedQuota), logger.LogQuota(actualQuota), logger.LogQuota(quotaDelta)) - if clamp != nil { - model.RecordLogWithAdminInfo(task.UserId, model.LogTypeSystem, logContent, - map[string]interface{}{"quota_saturation": clamp.AuditMap()}) - } else { - model.RecordLog(task.UserId, model.LogTypeSystem, logContent) - } - } - } else if quotaDelta < 0 { - // 需要退还多扣的费用 - refundQuota := -quotaDelta - logger.LogInfo(ctx, fmt.Sprintf("视频任务 %s 预扣费后返还:%s(实际消耗:%s,预扣费:%s,tokens:%d)", - task.TaskID, - logger.LogQuota(refundQuota), - logger.LogQuota(actualQuota), - logger.LogQuota(preConsumedQuota), - taskResult.TotalTokens, - )) - if err := model.IncreaseUserQuota(task.UserId, refundQuota, false); err != nil { - logger.LogError(ctx, fmt.Sprintf("退还预扣费失败: %s", err.Error())) - } else { - task.Quota = actualQuota // 更新任务记录的实际扣费额度 - - // 记录退款日志 - logContent := fmt.Sprintf("视频任务成功退还多扣费用,模型倍率 %.2f,分组倍率 %.2f,tokens %d,预扣费 %s,实际扣费 %s,退还 %s", - modelRatio, finalGroupRatio, taskResult.TotalTokens, - logger.LogQuota(preConsumedQuota), logger.LogQuota(actualQuota), logger.LogQuota(refundQuota)) - if clamp != nil { - model.RecordLogWithAdminInfo(task.UserId, model.LogTypeSystem, logContent, - map[string]interface{}{"quota_saturation": clamp.AuditMap()}) - } else { - model.RecordLog(task.UserId, model.LogTypeSystem, logContent) - } - } - } else { - // quotaDelta == 0, 预扣费刚好准确 - logger.LogInfo(ctx, fmt.Sprintf("视频任务 %s 预扣费准确(%s,tokens:%d)", - task.TaskID, logger.LogQuota(actualQuota), taskResult.TotalTokens)) - } - } - } - } - } - } - case model.TaskStatusFailure: - logger.LogJson(ctx, fmt.Sprintf("Task %s failed", taskId), task) - task.Status = model.TaskStatusFailure - task.Progress = "100%" - if task.FinishTime == 0 { - task.FinishTime = now - } - task.FailReason = taskResult.Reason - logger.LogInfo(ctx, fmt.Sprintf("Task %s failed: %s", task.TaskID, task.FailReason)) - taskResult.Progress = "100%" - if quota != 0 { - if preStatus != model.TaskStatusFailure { - shouldRefund = true - } else { - logger.LogWarn(ctx, fmt.Sprintf("Task %s already in failure status, skip refund", task.TaskID)) - } - } - default: - return fmt.Errorf("unknown task status %s for task %s", taskResult.Status, taskId) - } - if taskResult.Progress != "" { - task.Progress = taskResult.Progress - } - if err := task.Update(); err != nil { - common.SysLog("UpdateVideoTask task error: " + err.Error()) - shouldRefund = false - } - - if shouldRefund { - // 任务失败且之前状态不是失败才退还额度,防止重复退还 - if err := model.IncreaseUserQuota(task.UserId, quota, false); err != nil { - logger.LogWarn(ctx, "Failed to increase user quota: "+err.Error()) - } - logContent := fmt.Sprintf("Video async task failed %s, refund %s", task.TaskID, logger.LogQuota(quota)) - model.RecordLog(task.UserId, model.LogTypeSystem, logContent) - } - - return nil -} - -func redactVideoResponseBody(body []byte) []byte { - var m map[string]any - if err := json.Unmarshal(body, &m); err != nil { - return body - } - resp, _ := m["response"].(map[string]any) - if resp != nil { - delete(resp, "bytesBase64Encoded") - if v, ok := resp["video"].(string); ok { - resp["video"] = truncateBase64(v) - } - if vs, ok := resp["videos"].([]any); ok { - for i := range vs { - if vm, ok := vs[i].(map[string]any); ok { - delete(vm, "bytesBase64Encoded") - } - } - } - } - b, err := json.Marshal(m) - if err != nil { - return body - } - return b -} - -func truncateBase64(s string) string { - const maxKeep = 256 - if len(s) <= maxKeep { - return s - } - return s[:maxKeep] + "..." -} diff --git a/service/pre_consume_quota.go b/service/pre_consume_quota.go deleted file mode 100644 index d1c83656..00000000 --- a/service/pre_consume_quota.go +++ /dev/null @@ -1,79 +0,0 @@ -package service - -import ( - "fmt" - "net/http" - - "github.com/QuantumNous/new-api/common" - "github.com/QuantumNous/new-api/logger" - "github.com/QuantumNous/new-api/model" - relaycommon "github.com/QuantumNous/new-api/relay/common" - "github.com/QuantumNous/new-api/types" - - "github.com/bytedance/gopkg/util/gopool" - "github.com/gin-gonic/gin" -) - -func ReturnPreConsumedQuota(c *gin.Context, relayInfo *relaycommon.RelayInfo) { - if relayInfo.FinalPreConsumedQuota != 0 { - logger.LogInfo(c, fmt.Sprintf("用户 %d 请求失败, 返还预扣费额度 %s", relayInfo.UserId, logger.FormatQuota(relayInfo.FinalPreConsumedQuota))) - gopool.Go(func() { - relayInfoCopy := *relayInfo - - err := PostConsumeQuota(&relayInfoCopy, -relayInfoCopy.FinalPreConsumedQuota, 0, false) - if err != nil { - common.SysLog("error return pre-consumed quota: " + err.Error()) - } - }) - } -} - -// PreConsumeQuota checks if the user has enough quota to pre-consume. -// It returns the pre-consumed quota if successful, or an error if not. -func PreConsumeQuota(c *gin.Context, preConsumedQuota int, relayInfo *relaycommon.RelayInfo) *types.NewAPIError { - userQuota, err := model.GetUserQuota(relayInfo.UserId, false) - if err != nil { - return types.NewError(err, types.ErrorCodeQueryDataError, types.ErrOptionWithSkipRetry()) - } - if userQuota <= 0 { - return types.NewErrorWithStatusCode(fmt.Errorf("用户额度不足, 剩余额度: %s", logger.FormatQuota(userQuota)), types.ErrorCodeInsufficientUserQuota, http.StatusForbidden, types.ErrOptionWithSkipRetry(), types.ErrOptionWithNoRecordErrorLog()) - } - if userQuota-preConsumedQuota < 0 { - return types.NewErrorWithStatusCode(fmt.Errorf("预扣费额度失败, 用户剩余额度: %s, 需要预扣费额度: %s", logger.FormatQuota(userQuota), logger.FormatQuota(preConsumedQuota)), types.ErrorCodeInsufficientUserQuota, http.StatusForbidden, types.ErrOptionWithSkipRetry(), types.ErrOptionWithNoRecordErrorLog()) - } - - trustQuota := common.GetTrustQuota() - - relayInfo.UserQuota = userQuota - if userQuota > trustQuota { - // 用户额度充足,判断令牌额度是否充足 - if !relayInfo.TokenUnlimited { - // 非无限令牌,判断令牌额度是否充足 - tokenQuota := c.GetInt("token_quota") - if tokenQuota > trustQuota { - // 令牌额度充足,信任令牌 - preConsumedQuota = 0 - logger.LogInfo(c, fmt.Sprintf("用户 %d 剩余额度 %s 且令牌 %d 额度 %d 充足, 信任且不需要预扣费", relayInfo.UserId, logger.FormatQuota(userQuota), relayInfo.TokenId, tokenQuota)) - } - } else { - // in this case, we do not pre-consume quota - // because the user has enough quota - preConsumedQuota = 0 - logger.LogInfo(c, fmt.Sprintf("用户 %d 额度充足且为无限额度令牌, 信任且不需要预扣费", relayInfo.UserId)) - } - } - - if preConsumedQuota > 0 { - err := PreConsumeTokenQuota(relayInfo, preConsumedQuota) - if err != nil { - return types.NewErrorWithStatusCode(err, types.ErrorCodePreConsumeTokenQuotaFailed, http.StatusForbidden, types.ErrOptionWithSkipRetry(), types.ErrOptionWithNoRecordErrorLog()) - } - err = model.DecreaseUserQuota(relayInfo.UserId, preConsumedQuota, false) - if err != nil { - return types.NewError(err, types.ErrorCodeUpdateDataError, types.ErrOptionWithSkipRetry()) - } - logger.LogInfo(c, fmt.Sprintf("用户 %d 预扣费 %s, 预扣费后剩余额度: %s", relayInfo.UserId, logger.FormatQuota(preConsumedQuota), logger.FormatQuota(userQuota-preConsumedQuota))) - } - relayInfo.FinalPreConsumedQuota = preConsumedQuota - return nil -}