mirror of
https://github.com/yangjian102621/geekai.git
synced 2026-08-13 03:00:59 +00:00
feat(release): migrate GeekAI v4.3.0 to open source
- Sync backend and frontend from GeekAI Plus v4.3.0 - Remove commercial License flows and update open-source deployment defaults - Preserve Docker Compose deployment and bump image tags to v4.3.0 BREAKING CHANGE: commercial License configuration and related endpoints are removed
This commit is contained in:
@@ -14,7 +14,7 @@ import (
|
||||
"geekai/core/middleware"
|
||||
"geekai/core/types"
|
||||
"geekai/handler"
|
||||
logger2 "geekai/logger"
|
||||
"geekai/log"
|
||||
"geekai/service"
|
||||
"geekai/store/model"
|
||||
"geekai/store/vo"
|
||||
@@ -29,7 +29,7 @@ import (
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
var logger = logger2.GetLogger()
|
||||
var logger = log.GetLogger()
|
||||
|
||||
const SuperUsername = "admin"
|
||||
|
||||
@@ -293,7 +293,7 @@ func (h *ManagerHandler) ResetPass(c *gin.Context) {
|
||||
|
||||
password := utils.GenPassword(data.Password, user.Salt)
|
||||
user.Password = password
|
||||
res = h.DB.Updates(&user)
|
||||
res = h.DB.Model(&model.AdminUser{}).Where("id", data.Id).UpdateColumn("password", password)
|
||||
if res.Error != nil {
|
||||
resp.ERROR(c, res.Error.Error())
|
||||
return
|
||||
|
||||
@@ -8,7 +8,6 @@ package admin
|
||||
// * +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"geekai/core"
|
||||
"geekai/core/middleware"
|
||||
"geekai/core/types"
|
||||
@@ -59,15 +58,14 @@ func (h *ChatAppHandler) Save(c *gin.Context) {
|
||||
resp.ERROR(c, types.InvalidArgs)
|
||||
return
|
||||
}
|
||||
// 管理后台创建/编辑的 Gem 均为系统内置
|
||||
role.UserId = 0
|
||||
if data.SystemPrompt != "" {
|
||||
role.SystemPrompt = data.SystemPrompt
|
||||
}
|
||||
role.Id = data.Id
|
||||
if data.CreatedAt > 0 {
|
||||
role.CreatedAt = time.Unix(data.CreatedAt, 0)
|
||||
} else {
|
||||
err = h.DB.Where("marker", data.Key).First(&role).Error
|
||||
if err == nil {
|
||||
resp.ERROR(c, fmt.Sprintf("角色 %s 已存在", data.Key))
|
||||
return
|
||||
}
|
||||
}
|
||||
err = h.DB.Save(&role).Error
|
||||
if err != nil {
|
||||
@@ -83,7 +81,8 @@ func (h *ChatAppHandler) Save(c *gin.Context) {
|
||||
func (h *ChatAppHandler) List(c *gin.Context) {
|
||||
var items []model.ChatApp
|
||||
var roles = make([]vo.ChatApp, 0)
|
||||
res := h.DB.Order("sort_num ASC").Find(&items)
|
||||
// 仅展示系统内置智能体(user_id = 0),不展示用户自行创建的智能体
|
||||
res := h.DB.Where("user_id = 0").Order("sort_num ASC").Find(&items)
|
||||
if res.Error != nil {
|
||||
resp.ERROR(c, "No data found")
|
||||
return
|
||||
|
||||
@@ -8,6 +8,7 @@ package admin
|
||||
// * +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"geekai/core"
|
||||
"geekai/core/middleware"
|
||||
"geekai/core/types"
|
||||
@@ -35,6 +36,7 @@ type ConfigHandler struct {
|
||||
smtpService *service.SmtpService
|
||||
captchaService *service.CaptchaService
|
||||
wxLoginService *service.WxLoginService
|
||||
wxGzhService *service.WxGzhService
|
||||
}
|
||||
|
||||
func NewConfigHandler(
|
||||
@@ -49,6 +51,7 @@ func NewConfigHandler(
|
||||
smtpService *service.SmtpService,
|
||||
captchaService *service.CaptchaService,
|
||||
wxLoginService *service.WxLoginService,
|
||||
wxChatService *service.WxGzhService,
|
||||
) *ConfigHandler {
|
||||
return &ConfigHandler{
|
||||
BaseHandler: handler.BaseHandler{App: app, DB: db},
|
||||
@@ -61,6 +64,7 @@ func NewConfigHandler(
|
||||
smtpService: smtpService,
|
||||
captchaService: captchaService,
|
||||
wxLoginService: wxLoginService,
|
||||
wxGzhService: wxChatService,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -84,6 +88,7 @@ func (h *ConfigHandler) RegisterRoutes() {
|
||||
rg.POST("update/oss", h.UpdateOss)
|
||||
rg.POST("update/smtp", h.UpdateStmp)
|
||||
rg.GET("get", h.Get)
|
||||
rg.POST("update/wx_gzh", h.UpdateWxGzh)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -110,15 +115,18 @@ func (h *ConfigHandler) UpdateBase(c *gin.Context) {
|
||||
// UpdatePower 更新系统配置
|
||||
func (h *ConfigHandler) UpdatePower(c *gin.Context) {
|
||||
var data struct {
|
||||
InitPower int `json:"init_power,omitempty"` // 新用户注册赠送算力值
|
||||
DailyPower int `json:"daily_power,omitempty"` // 每日签到赠送算力
|
||||
InvitePower int `json:"invite_power,omitempty"` // 邀请新用户赠送算力值
|
||||
MjPower int `json:"mj_power,omitempty"` // MJ 绘画消耗算力
|
||||
MjActionPower int `json:"mj_action_power,omitempty"` // MJ 操作(放大,变换)消耗算力
|
||||
SdPower int `json:"sd_power,omitempty"` // SD 绘画消耗算力
|
||||
SunoPower int `json:"suno_power,omitempty"` // Suno 生成歌曲消耗算力
|
||||
LumaPower int `json:"luma_power,omitempty"` // Luma 生成视频消耗算力
|
||||
KeLingPowers map[string]int `json:"keling_powers,omitempty"` // 可灵生成视频消耗算力
|
||||
InitPower int `json:"init_power,omitempty"` // 新用户注册赠送算力值
|
||||
DailyPower int `json:"daily_power,omitempty"` // 每日签到赠送算力
|
||||
InvitePower int `json:"invite_power,omitempty"` // 邀请新用户赠送算力值
|
||||
MjPower int `json:"mj_power,omitempty"` // MJ 绘画消耗算力
|
||||
MjActionPower int `json:"mj_action_power,omitempty"` // MJ 操作(放大,变换)消耗算力
|
||||
MjUpscalePower int `json:"mj_upscale_power,omitempty"` // MJ 放大/变换消耗算力
|
||||
MjBlendPower int `json:"mj_blend_power,omitempty"` // MJ 融图消耗算力
|
||||
MjSwapFacePower int `json:"mj_swap_face_power,omitempty"` // MJ 换脸消耗算力
|
||||
MjModalPower int `json:"mj_modal_power,omitempty"` // MJ 局部重绘消耗算力
|
||||
SunoPower int `json:"suno_power,omitempty"` // Suno 生成歌曲消耗算力
|
||||
LumaPower int `json:"luma_power,omitempty"` // Luma 生成视频消耗算力
|
||||
KeLingPowers map[string]int `json:"keling_powers,omitempty"` // 可灵生成视频消耗算力
|
||||
}
|
||||
if err := c.ShouldBindJSON(&data); err != nil {
|
||||
resp.ERROR(c, types.InvalidArgs)
|
||||
@@ -130,7 +138,10 @@ func (h *ConfigHandler) UpdatePower(c *gin.Context) {
|
||||
h.sysConfig.Base.InvitePower = data.InvitePower
|
||||
h.sysConfig.Base.MjPower = data.MjPower
|
||||
h.sysConfig.Base.MjActionPower = data.MjActionPower
|
||||
h.sysConfig.Base.SdPower = data.SdPower
|
||||
h.sysConfig.Base.MjUpscalePower = data.MjUpscalePower
|
||||
h.sysConfig.Base.MjBlendPower = data.MjBlendPower
|
||||
h.sysConfig.Base.MjSwapFacePower = data.MjSwapFacePower
|
||||
h.sysConfig.Base.MjModalPower = data.MjModalPower
|
||||
h.sysConfig.Base.SunoPower = data.SunoPower
|
||||
h.sysConfig.Base.LumaPower = data.LumaPower
|
||||
h.sysConfig.Base.KeLingPowers = data.KeLingPowers
|
||||
@@ -374,14 +385,19 @@ func (h *ConfigHandler) Update(name string, value any) error {
|
||||
func (h *ConfigHandler) Get(c *gin.Context) {
|
||||
name := c.Query("key")
|
||||
var config model.Config
|
||||
res := h.DB.Where("name", name).First(&config)
|
||||
if res.Error != nil {
|
||||
resp.ERROR(c, res.Error.Error())
|
||||
err := h.DB.Where("name", name).First(&config).Error
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
resp.SUCCESS(c, nil)
|
||||
return
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
resp.ERROR(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
var value map[string]any
|
||||
err := utils.JsonDecode(config.Value, &value)
|
||||
err = utils.JsonDecode(config.Value, &value)
|
||||
if err != nil {
|
||||
resp.ERROR(c, err.Error())
|
||||
return
|
||||
@@ -389,3 +405,20 @@ func (h *ConfigHandler) Get(c *gin.Context) {
|
||||
|
||||
resp.SUCCESS(c, value)
|
||||
}
|
||||
|
||||
func (h *ConfigHandler) UpdateWxGzh(c *gin.Context) {
|
||||
var data types.WxGzhConfig
|
||||
if err := c.ShouldBindJSON(&data); err != nil {
|
||||
resp.ERROR(c, types.InvalidArgs)
|
||||
return
|
||||
}
|
||||
err := h.Update(types.ConfigKeyWxGzh, data)
|
||||
if err != nil {
|
||||
resp.ERROR(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
h.wxGzhService.UpdateConfig(data)
|
||||
h.sysConfig.WxGzh = data
|
||||
resp.SUCCESS(c, data)
|
||||
}
|
||||
|
||||
@@ -123,22 +123,20 @@ func (h *DashboardHandler) Stats(c *gin.Context) {
|
||||
h.DB.Model(&model.Order{}).Where("status = ?", types.OrderPaidSuccess).Where("created_at > ?", zeroTime).Count(&stats.TodayOrders)
|
||||
|
||||
// 图片生成任务统计
|
||||
var mjJobs, sdJobs, dallJobs, jimengImageJobs int64
|
||||
var mjJobs, imageJobs, jimengImageJobs int64
|
||||
h.DB.Model(&model.MidJourneyJob{}).Count(&mjJobs)
|
||||
h.DB.Model(&model.SdJob{}).Count(&sdJobs)
|
||||
h.DB.Model(&model.DallJob{}).Count(&dallJobs)
|
||||
h.DB.Model(&model.ImageJob{}).Count(&imageJobs)
|
||||
h.DB.Model(&model.JimengJob{}).Where("type IN ?", []string{"text_to_image", "image_to_image", "image_edit", "image_effects"}).Count(&jimengImageJobs)
|
||||
stats.ImageJobs = mjJobs + sdJobs + dallJobs + jimengImageJobs
|
||||
stats.ImageJobs = mjJobs + imageJobs + jimengImageJobs
|
||||
|
||||
logger.Info("stats.ImageJobs", stats.ImageJobs)
|
||||
|
||||
// 今日图片生成任务统计
|
||||
var todayMjJobs, todaySdJobs, todayDallJobs, todayJimengImageJobs int64
|
||||
var todayMjJobs, todayImageJobs, todayJimengImageJobs int64
|
||||
h.DB.Model(&model.MidJourneyJob{}).Where("created_at > ?", zeroTime).Count(&todayMjJobs)
|
||||
h.DB.Model(&model.SdJob{}).Where("created_at > ?", zeroTime).Count(&todaySdJobs)
|
||||
h.DB.Model(&model.DallJob{}).Where("created_at > ?", zeroTime).Count(&todayDallJobs)
|
||||
h.DB.Model(&model.ImageJob{}).Where("created_at > ?", zeroTime).Count(&todayImageJobs)
|
||||
h.DB.Model(&model.JimengJob{}).Where("type IN ?", []string{"text_to_image", "image_to_image", "image_edit", "image_effects"}).Where("created_at > ?", zeroTime).Count(&todayJimengImageJobs)
|
||||
stats.TodayImageJobs = todayMjJobs + todaySdJobs + todayDallJobs + todayJimengImageJobs
|
||||
stats.TodayImageJobs = todayMjJobs + todayImageJobs + todayJimengImageJobs
|
||||
|
||||
// 视频生成任务统计
|
||||
var videoJobs, jimengVideoJobs int64
|
||||
|
||||
@@ -42,8 +42,7 @@ func (h *ImageHandler) RegisterRoutes() {
|
||||
group.Use(middleware.AdminAuthMiddleware(h.App.Config.AdminSession.SecretKey, h.App.Redis))
|
||||
{
|
||||
group.POST("list/mj", h.MjList)
|
||||
group.POST("list/sd", h.SdList)
|
||||
group.POST("list/dall", h.DallList)
|
||||
group.POST("list/image", h.ImageList)
|
||||
group.GET("remove", h.Remove)
|
||||
}
|
||||
}
|
||||
@@ -100,8 +99,8 @@ func (h *ImageHandler) MjList(c *gin.Context) {
|
||||
resp.SUCCESS(c, vo.NewPage(total, data.Page, data.PageSize, items))
|
||||
}
|
||||
|
||||
// SdList Stable Diffusion 任务列表
|
||||
func (h *ImageHandler) SdList(c *gin.Context) {
|
||||
// ImageList Image generation 任务列表
|
||||
func (h *ImageHandler) ImageList(c *gin.Context) {
|
||||
var data imageQuery
|
||||
if err := c.ShouldBindJSON(&data); err != nil {
|
||||
resp.ERROR(c, types.InvalidArgs)
|
||||
@@ -123,59 +122,15 @@ func (h *ImageHandler) SdList(c *gin.Context) {
|
||||
session = session.Where("created_at >= ? AND created_at <= ?", data.CreatedAt[0], data.CreatedAt[1])
|
||||
}
|
||||
var total int64
|
||||
session.Model(&model.SdJob{}).Count(&total)
|
||||
var list []model.SdJob
|
||||
var items = make([]vo.SdJob, 0)
|
||||
session.Model(&model.ImageJob{}).Count(&total)
|
||||
var list []model.ImageJob
|
||||
var items = make([]vo.ImageJob, 0)
|
||||
offset := (data.Page - 1) * data.PageSize
|
||||
err := session.Order("id DESC").Offset(offset).Limit(data.PageSize).Find(&list).Error
|
||||
if err == nil {
|
||||
// 填充数据
|
||||
for _, item := range list {
|
||||
var job vo.SdJob
|
||||
err = utils.CopyObject(item, &job)
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
job.CreatedAt = item.CreatedAt.Unix()
|
||||
items = append(items, job)
|
||||
}
|
||||
}
|
||||
|
||||
resp.SUCCESS(c, vo.NewPage(total, data.Page, data.PageSize, items))
|
||||
}
|
||||
|
||||
// DallList DALL-E 任务列表
|
||||
func (h *ImageHandler) DallList(c *gin.Context) {
|
||||
var data imageQuery
|
||||
if err := c.ShouldBindJSON(&data); err != nil {
|
||||
resp.ERROR(c, types.InvalidArgs)
|
||||
return
|
||||
}
|
||||
|
||||
session := h.DB.Session(&gorm.Session{})
|
||||
if data.Username != "" {
|
||||
var user model.User
|
||||
err := h.DB.Where("username", data.Username).First(&user).Error
|
||||
if err == nil {
|
||||
session = session.Where("user_id", user.Id)
|
||||
}
|
||||
}
|
||||
if data.Prompt != "" {
|
||||
session = session.Where("prompt LIKE ?", "%"+data.Prompt+"%")
|
||||
}
|
||||
if len(data.CreatedAt) == 2 {
|
||||
session = session.Where("created_at >= ? AND created_at <= ?", data.CreatedAt[0], data.CreatedAt[1])
|
||||
}
|
||||
var total int64
|
||||
session.Model(&model.DallJob{}).Count(&total)
|
||||
var list []model.DallJob
|
||||
var items = make([]vo.DallJob, 0)
|
||||
offset := (data.Page - 1) * data.PageSize
|
||||
err := session.Order("id DESC").Offset(offset).Limit(data.PageSize).Find(&list).Error
|
||||
if err == nil {
|
||||
// 填充数据
|
||||
for _, item := range list {
|
||||
var job vo.DallJob
|
||||
var job vo.ImageJob
|
||||
err = utils.CopyObject(item, &job)
|
||||
if err != nil {
|
||||
continue
|
||||
@@ -209,8 +164,8 @@ func (h *ImageHandler) Remove(c *gin.Context) {
|
||||
remark = fmt.Sprintf("任务失败,退回算力。任务ID:%d,Err: %s", job.Id, job.ErrMsg)
|
||||
progress = job.Progress
|
||||
imgURL = job.ImgURL
|
||||
case "sd":
|
||||
var job model.SdJob
|
||||
case "image":
|
||||
var job model.ImageJob
|
||||
if res := h.DB.Where("id", id).First(&job); res.Error != nil {
|
||||
resp.ERROR(c, "记录不存在")
|
||||
return
|
||||
@@ -218,22 +173,7 @@ func (h *ImageHandler) Remove(c *gin.Context) {
|
||||
|
||||
// 删除任务
|
||||
tx.Delete(&job)
|
||||
md = "stable-diffusion"
|
||||
power = job.Power
|
||||
userId = int(job.UserId)
|
||||
remark = fmt.Sprintf("任务失败,退回算力。任务ID:%d,Err: %s", job.Id, job.ErrMsg)
|
||||
progress = job.Progress
|
||||
imgURL = job.ImgURL
|
||||
case "dall":
|
||||
var job model.DallJob
|
||||
if res := h.DB.Where("id", id).First(&job); res.Error != nil {
|
||||
resp.ERROR(c, "记录不存在")
|
||||
return
|
||||
}
|
||||
|
||||
// 删除任务
|
||||
tx.Delete(&job)
|
||||
md = "dall-e-3"
|
||||
md = "image-generation"
|
||||
power = job.Power
|
||||
userId = int(job.UserId)
|
||||
remark = fmt.Sprintf("任务失败,退回算力。任务ID:%d,Err: %s", job.Id, job.ErrMsg)
|
||||
|
||||
@@ -1,215 +0,0 @@
|
||||
package admin
|
||||
|
||||
// * +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
|
||||
// * Copyright 2023 The Geek-AI Authors. All rights reserved.
|
||||
// * Use of this source code is governed by a Apache-2.0 license
|
||||
// * that can be found in the LICENSE file.
|
||||
// * @Author yangjian102621@163.com
|
||||
// * +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"geekai/core"
|
||||
"geekai/core/middleware"
|
||||
"geekai/core/types"
|
||||
"geekai/handler"
|
||||
"geekai/service"
|
||||
"geekai/service/oss"
|
||||
"geekai/store/model"
|
||||
"geekai/store/vo"
|
||||
"geekai/utils"
|
||||
"geekai/utils/resp"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type MediaHandler struct {
|
||||
handler.BaseHandler
|
||||
userService *service.UserService
|
||||
uploader *oss.UploaderManager
|
||||
}
|
||||
|
||||
func NewMediaHandler(app *core.AppServer, db *gorm.DB, userService *service.UserService, manager *oss.UploaderManager) *MediaHandler {
|
||||
return &MediaHandler{BaseHandler: handler.BaseHandler{App: app, DB: db}, userService: userService, uploader: manager}
|
||||
}
|
||||
|
||||
// RegisterRoutes 注册路由
|
||||
func (h *MediaHandler) RegisterRoutes() {
|
||||
group := h.App.Engine.Group("/api/admin/media/")
|
||||
|
||||
// 需要管理员授权的接口
|
||||
group.Use(middleware.AdminAuthMiddleware(h.App.Config.AdminSession.SecretKey, h.App.Redis))
|
||||
{
|
||||
group.POST("suno", h.SunoList)
|
||||
group.POST("videos", h.Videos)
|
||||
group.GET("remove", h.Remove)
|
||||
}
|
||||
}
|
||||
|
||||
type mediaQuery struct {
|
||||
Type string `json:"type"` // 任务类型 luma, keling
|
||||
Prompt string `json:"prompt"`
|
||||
Username string `json:"username"`
|
||||
CreatedAt []string `json:"created_at"`
|
||||
Page int `json:"page"`
|
||||
PageSize int `json:"page_size"`
|
||||
}
|
||||
|
||||
// SunoList Suno 任务列表
|
||||
func (h *MediaHandler) SunoList(c *gin.Context) {
|
||||
var data mediaQuery
|
||||
if err := c.ShouldBindJSON(&data); err != nil {
|
||||
resp.ERROR(c, types.InvalidArgs)
|
||||
return
|
||||
}
|
||||
|
||||
session := h.DB.Session(&gorm.Session{})
|
||||
if data.Username != "" {
|
||||
var user model.User
|
||||
err := h.DB.Where("username", data.Username).First(&user).Error
|
||||
if err == nil {
|
||||
session = session.Where("user_id", user.Id)
|
||||
}
|
||||
}
|
||||
if data.Prompt != "" {
|
||||
session = session.Where("prompt LIKE ?", "%"+data.Prompt+"%")
|
||||
}
|
||||
if len(data.CreatedAt) == 2 {
|
||||
session = session.Where("created_at >= ? AND created_at <= ?", data.CreatedAt[0], data.CreatedAt[1])
|
||||
}
|
||||
var total int64
|
||||
session.Model(&model.SunoJob{}).Count(&total)
|
||||
var list []model.SunoJob
|
||||
var items = make([]vo.SunoJob, 0)
|
||||
offset := (data.Page - 1) * data.PageSize
|
||||
err := session.Order("id DESC").Offset(offset).Limit(data.PageSize).Find(&list).Error
|
||||
if err == nil {
|
||||
// 填充数据
|
||||
for _, item := range list {
|
||||
var job vo.SunoJob
|
||||
err = utils.CopyObject(item, &job)
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
job.CreatedAt = item.CreatedAt.Unix()
|
||||
items = append(items, job)
|
||||
}
|
||||
}
|
||||
|
||||
resp.SUCCESS(c, vo.NewPage(total, data.Page, data.PageSize, items))
|
||||
}
|
||||
|
||||
// Videos 视频任务列表
|
||||
func (h *MediaHandler) Videos(c *gin.Context) {
|
||||
var data mediaQuery
|
||||
if err := c.ShouldBindJSON(&data); err != nil {
|
||||
resp.ERROR(c, types.InvalidArgs)
|
||||
return
|
||||
}
|
||||
|
||||
session := h.DB.Session(&gorm.Session{}).Where("type", data.Type)
|
||||
if data.Username != "" {
|
||||
var user model.User
|
||||
err := h.DB.Where("username", data.Username).First(&user).Error
|
||||
if err == nil {
|
||||
session = session.Where("user_id", user.Id)
|
||||
}
|
||||
}
|
||||
if data.Prompt != "" {
|
||||
session = session.Where("prompt LIKE ?", "%"+data.Prompt+"%")
|
||||
}
|
||||
if len(data.CreatedAt) == 2 {
|
||||
session = session.Where("created_at >= ? AND created_at <= ?", data.CreatedAt[0], data.CreatedAt[1])
|
||||
}
|
||||
var total int64
|
||||
session.Model(&model.VideoJob{}).Count(&total)
|
||||
var list []model.VideoJob
|
||||
var items = make([]vo.VideoJob, 0)
|
||||
offset := (data.Page - 1) * data.PageSize
|
||||
err := session.Order("id DESC").Offset(offset).Limit(data.PageSize).Find(&list).Error
|
||||
if err == nil {
|
||||
// 填充数据
|
||||
for _, item := range list {
|
||||
var job vo.VideoJob
|
||||
err = utils.CopyObject(item, &job)
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
job.CreatedAt = item.CreatedAt.Unix()
|
||||
if job.VideoURL == "" {
|
||||
job.VideoURL = job.WaterURL
|
||||
}
|
||||
items = append(items, job)
|
||||
}
|
||||
}
|
||||
|
||||
resp.SUCCESS(c, vo.NewPage(total, data.Page, data.PageSize, items))
|
||||
}
|
||||
|
||||
func (h *MediaHandler) Remove(c *gin.Context) {
|
||||
id := h.GetInt(c, "id", 0)
|
||||
tab := c.Query("tab")
|
||||
|
||||
tx := h.DB.Begin()
|
||||
var md, remark, fileURL string
|
||||
var power, userId, progress int
|
||||
switch tab {
|
||||
case "suno":
|
||||
var job model.SunoJob
|
||||
if err := h.DB.Where("id", id).First(&job).Error; err != nil {
|
||||
resp.ERROR(c, "记录不存在")
|
||||
return
|
||||
}
|
||||
tx.Delete(&job)
|
||||
md = "suno"
|
||||
power = job.Power
|
||||
userId = int(job.UserId)
|
||||
remark = fmt.Sprintf("SUNO 任务失败,退回算力。任务ID:%d,Err: %s", job.Id, job.ErrMsg)
|
||||
progress = job.Progress
|
||||
fileURL = job.AudioURL
|
||||
case "luma":
|
||||
case "keling":
|
||||
var job model.VideoJob
|
||||
if res := h.DB.Where("id", id).First(&job); res.Error != nil {
|
||||
resp.ERROR(c, "记录不存在")
|
||||
return
|
||||
}
|
||||
|
||||
// 删除任务
|
||||
tx.Delete(&job)
|
||||
md = job.Type
|
||||
power = job.Power
|
||||
userId = int(job.UserId)
|
||||
remark = fmt.Sprintf("LUMA 任务失败,退回算力。任务ID:%d,Err: %s", job.Id, job.ErrMsg)
|
||||
progress = job.Progress
|
||||
fileURL = job.VideoURL
|
||||
if fileURL == "" {
|
||||
fileURL = job.WaterURL
|
||||
}
|
||||
default:
|
||||
resp.ERROR(c, types.InvalidArgs)
|
||||
return
|
||||
}
|
||||
|
||||
if progress != 100 {
|
||||
err := h.userService.IncreasePower(uint(userId), power, model.PowerLog{
|
||||
Type: types.PowerRefund,
|
||||
Model: md,
|
||||
Remark: remark,
|
||||
})
|
||||
if err != nil {
|
||||
tx.Rollback()
|
||||
resp.ERROR(c, err.Error())
|
||||
return
|
||||
}
|
||||
}
|
||||
tx.Commit()
|
||||
// remove image
|
||||
err := h.uploader.GetUploadHandler().Delete(fileURL)
|
||||
if err != nil {
|
||||
logger.Error("remove image failed: ", err)
|
||||
}
|
||||
|
||||
resp.SUCCESS(c)
|
||||
}
|
||||
@@ -212,12 +212,8 @@ func (h *ModerationHandler) GetSourceList(c *gin.Context) {
|
||||
"name": "Midjourney 绘图",
|
||||
},
|
||||
{
|
||||
"id": types.ModerationSourceDalle,
|
||||
"name": "Dalle 绘图",
|
||||
},
|
||||
{
|
||||
"id": types.ModerationSourceSD,
|
||||
"name": "StableDiffusion 绘图",
|
||||
"id": types.ModerationSourceImage,
|
||||
"name": "AI图像生成",
|
||||
},
|
||||
{
|
||||
"id": types.ModerationSourceSuno,
|
||||
|
||||
@@ -0,0 +1,284 @@
|
||||
package admin
|
||||
|
||||
// * +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
|
||||
// * Copyright 2023 The Geek-AI Authors. All rights reserved.
|
||||
// * Use of this source code is governed by a Apache-2.0 license
|
||||
// * that can be found in the LICENSE file.
|
||||
// * @Author yangjian102621@163.com
|
||||
// * +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
|
||||
|
||||
import (
|
||||
"geekai/core"
|
||||
"geekai/core/middleware"
|
||||
"geekai/core/types"
|
||||
"geekai/handler"
|
||||
"geekai/service/ppt"
|
||||
"geekai/store/model"
|
||||
"geekai/utils"
|
||||
"geekai/utils/resp"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// PPTHandler 管理后台 PPT 生成配置处理器
|
||||
type PPTHandler struct {
|
||||
handler.BaseHandler
|
||||
pptService *ppt.PptService
|
||||
}
|
||||
|
||||
// NewPPTHandler 创建管理后台 PPT 配置处理器
|
||||
func NewPPTHandler(app *core.AppServer, db *gorm.DB, pptService *ppt.PptService) *PPTHandler {
|
||||
return &PPTHandler{
|
||||
BaseHandler: handler.BaseHandler{App: app, DB: db},
|
||||
pptService: pptService,
|
||||
}
|
||||
}
|
||||
|
||||
// RegisterRoutes 注册 PPT 配置相关路由
|
||||
func (h *PPTHandler) RegisterRoutes() {
|
||||
rg := h.App.Engine.Group("/api/admin/ppt/")
|
||||
rg.Use(middleware.AdminAuthMiddleware(h.App.Config.AdminSession.SecretKey, h.App.Redis))
|
||||
{
|
||||
rg.GET("config", h.GetConfig)
|
||||
rg.POST("config/update", h.UpdateConfig)
|
||||
rg.GET("jobs", h.Jobs)
|
||||
rg.GET("jobs/:task_id", h.JobDetail)
|
||||
rg.GET("jobs/:task_id/export", h.ExportJob)
|
||||
rg.GET("stats", h.Stats)
|
||||
}
|
||||
}
|
||||
|
||||
// GetConfig 获取 PPT 生成配置
|
||||
func (h *PPTHandler) GetConfig(c *gin.Context) {
|
||||
var cfg model.Config
|
||||
err := h.DB.Where("name", types.ConfigKeyPPT).First(&cfg).Error
|
||||
if err != nil {
|
||||
if err == gorm.ErrRecordNotFound {
|
||||
// 返回一个默认空配置
|
||||
resp.SUCCESS(c, types.PPTConfig{
|
||||
OutlineLLMModel: "gpt-4o-mini",
|
||||
MaxSlidesPerTask: 30,
|
||||
PowerCostPerSlide: 0,
|
||||
MaxConcurrentRequests: 3,
|
||||
QPSLimit: 1,
|
||||
NanoBananaModel: "nano-banana",
|
||||
NanoBananaAspectRatio: "16:9",
|
||||
SeedreamSize: "1920x1080",
|
||||
})
|
||||
return
|
||||
}
|
||||
resp.ERROR(c, "获取配置失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
var pptConfig types.PPTConfig
|
||||
err = utils.JsonDecode(cfg.Value, &pptConfig)
|
||||
if err != nil {
|
||||
resp.ERROR(c, "解析配置失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
resp.SUCCESS(c, pptConfig)
|
||||
}
|
||||
|
||||
// UpdateConfig 更新 PPT 生成配置
|
||||
func (h *PPTHandler) UpdateConfig(c *gin.Context) {
|
||||
var req types.PPTConfig
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
resp.ERROR(c, "参数错误")
|
||||
return
|
||||
}
|
||||
|
||||
// 基础校验
|
||||
if req.MaxSlidesPerTask <= 0 {
|
||||
resp.ERROR(c, "单个任务最多 PPT 页数必须大于 0")
|
||||
return
|
||||
}
|
||||
|
||||
if req.PowerCostPerSlide < 0 {
|
||||
resp.ERROR(c, "每张 PPT 图片消耗算力不能小于 0")
|
||||
return
|
||||
}
|
||||
|
||||
if req.MaxConcurrentRequests <= 0 {
|
||||
req.MaxConcurrentRequests = 3
|
||||
}
|
||||
if req.QPSLimit <= 0 {
|
||||
req.QPSLimit = 1
|
||||
}
|
||||
|
||||
// 根据当前生图提供方做必填校验
|
||||
switch req.ActiveImageProvider {
|
||||
case types.PPTImageProviderNanoBanana:
|
||||
if req.NanoBananaApiURL == "" {
|
||||
resp.ERROR(c, "Nano Banana API 地址不能为空")
|
||||
return
|
||||
}
|
||||
if req.NanoBananaApiKey == "" {
|
||||
resp.ERROR(c, "Nano Banana API Key 不能为空")
|
||||
return
|
||||
}
|
||||
case types.PPTImageProviderSeedream:
|
||||
if req.SeedreamBaseURL == "" {
|
||||
resp.ERROR(c, "Seedream Base URL 不能为空")
|
||||
return
|
||||
}
|
||||
if req.SeedreamApiKey == "" {
|
||||
resp.ERROR(c, "Seedream API Key 不能为空")
|
||||
return
|
||||
}
|
||||
if req.SeedreamModel == "" {
|
||||
resp.ERROR(c, "Seedream 模型 ID 不能为空")
|
||||
return
|
||||
}
|
||||
default:
|
||||
// 允许为空,未来可以扩展更多 provider
|
||||
}
|
||||
|
||||
value := utils.JsonEncode(&req)
|
||||
var cfg model.Config
|
||||
err := h.DB.Where("name", types.ConfigKeyPPT).First(&cfg).Error
|
||||
if err != nil {
|
||||
if err == gorm.ErrRecordNotFound {
|
||||
cfg.Name = types.ConfigKeyPPT
|
||||
cfg.Value = value
|
||||
if err = h.DB.Create(&cfg).Error; err != nil {
|
||||
resp.ERROR(c, "创建配置失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
resp.SUCCESS(c, gin.H{"message": "配置创建成功"})
|
||||
return
|
||||
}
|
||||
resp.ERROR(c, "获取配置失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
cfg.Value = value
|
||||
if err = h.DB.Updates(&cfg).Error; err != nil {
|
||||
resp.ERROR(c, "更新配置失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
resp.SUCCESS(c, gin.H{"message": "配置更新成功"})
|
||||
}
|
||||
|
||||
// Jobs 管理后台查看 PPT 任务列表(内存任务)
|
||||
func (h *PPTHandler) Jobs(c *gin.Context) {
|
||||
page := h.GetInt(c, "page", 1)
|
||||
pageSize := h.GetInt(c, "page_size", 20)
|
||||
filterUserId := h.GetInt(c, "user_id", 0)
|
||||
status := h.GetTrim(c, "status")
|
||||
|
||||
filtered, total := h.pptService.ListAdminJobs(c.Request.Context(), page, pageSize, filterUserId, status)
|
||||
|
||||
jobs := make([]gin.H, 0, len(filtered))
|
||||
for _, t := range filtered {
|
||||
job := t.TaskSummaryMap()
|
||||
job["user_id"] = t.UserID
|
||||
job["error_message"] = t.ErrorMessage
|
||||
jobs = append(jobs, job)
|
||||
}
|
||||
|
||||
resp.SUCCESS(c, gin.H{
|
||||
"jobs": jobs,
|
||||
"total": total,
|
||||
"page": page,
|
||||
"page_size": pageSize,
|
||||
})
|
||||
}
|
||||
|
||||
func buildAdminPPTTaskDetail(task *ppt.Task) gin.H {
|
||||
percentage := 0
|
||||
if task.Total > 0 {
|
||||
percentage = int(float64(task.Completed) / float64(task.Total) * 100)
|
||||
}
|
||||
|
||||
return gin.H{
|
||||
"task_id": task.TaskID,
|
||||
"user_id": task.UserID,
|
||||
"status": task.Status,
|
||||
"progress": gin.H{
|
||||
"total_slides": task.Total,
|
||||
"completed_slides": task.Completed,
|
||||
"percentage": percentage,
|
||||
},
|
||||
"slides": task.Slides,
|
||||
"error_message": task.ErrorMessage,
|
||||
"content": task.Content,
|
||||
"prompt": task.Prompt,
|
||||
"title": task.Title,
|
||||
"thumb": task.Thumb,
|
||||
}
|
||||
}
|
||||
|
||||
// JobDetail 管理后台查看指定 PPT 任务详情
|
||||
func (h *PPTHandler) JobDetail(c *gin.Context) {
|
||||
taskID := c.Param("task_id")
|
||||
if taskID == "" {
|
||||
resp.ERROR(c, types.InvalidArgs)
|
||||
return
|
||||
}
|
||||
|
||||
task, ok := h.pptService.GetTask(taskID)
|
||||
if !ok {
|
||||
resp.ERROR(c, "任务不存在")
|
||||
return
|
||||
}
|
||||
h.pptService.EnsureTaskMeta(c.Request.Context(), task)
|
||||
resp.SUCCESS(c, buildAdminPPTTaskDetail(task))
|
||||
}
|
||||
|
||||
// ExportJob 管理后台导出 PPT 任务
|
||||
func (h *PPTHandler) ExportJob(c *gin.Context) {
|
||||
taskID := c.Param("task_id")
|
||||
if taskID == "" {
|
||||
resp.ERROR(c, types.InvalidArgs)
|
||||
return
|
||||
}
|
||||
ef, ok := ppt.ParseExportFormat(c.Query("format"))
|
||||
if !ok {
|
||||
resp.ERROR(c, "format 参数无效,支持 pdf 或 pptx")
|
||||
return
|
||||
}
|
||||
|
||||
task, exists := h.pptService.GetTask(taskID)
|
||||
if !exists {
|
||||
resp.ERROR(c, "任务不存在")
|
||||
return
|
||||
}
|
||||
if task.Status != ppt.TaskStatusCompleted {
|
||||
resp.ERROR(c, "仅已完成任务可导出")
|
||||
return
|
||||
}
|
||||
|
||||
h.pptService.EnsureTaskMeta(c.Request.Context(), task)
|
||||
|
||||
ossCfg := types.OSSConfig{}
|
||||
if h.App.SysConfig != nil {
|
||||
ossCfg = h.App.SysConfig.OSS
|
||||
}
|
||||
data, err := ppt.BuildExportBytes(c.Request.Context(), task.Slides, ef, ossCfg, h.App.Config)
|
||||
if err != nil {
|
||||
resp.ERROR(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
base := ppt.SanitizeExportBaseName(task.Title, task.TaskID)
|
||||
filename := base + ppt.ExportFileExt(ef)
|
||||
c.Header("Content-Disposition", ppt.ContentDispositionAttachment(filename))
|
||||
c.Data(200, ppt.ExportMimeType(ef), data)
|
||||
}
|
||||
|
||||
// Stats PPT 任务统计信息
|
||||
func (h *PPTHandler) Stats(c *gin.Context) {
|
||||
total, completed, processing, failed, pending := h.pptService.Stats()
|
||||
|
||||
resp.SUCCESS(c, gin.H{
|
||||
"totalTasks": total,
|
||||
"completedTasks": completed,
|
||||
"processingTasks": processing,
|
||||
"failedTasks": failed,
|
||||
"pendingTasks": pending,
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,137 @@
|
||||
package admin
|
||||
|
||||
// * +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
|
||||
// * Copyright 2023 The Geek-AI Authors. All rights reserved.
|
||||
// * Use of this source code is governed by a Apache-2.0 license
|
||||
// * that can be found in the LICENSE file.
|
||||
// * @Author yangjian102621@163.com
|
||||
// * +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"geekai/core"
|
||||
"geekai/core/middleware"
|
||||
"geekai/core/types"
|
||||
"geekai/handler"
|
||||
"geekai/service"
|
||||
"geekai/service/oss"
|
||||
"geekai/store/model"
|
||||
"geekai/store/vo"
|
||||
"geekai/utils"
|
||||
"geekai/utils/resp"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type SunoHandler struct {
|
||||
handler.BaseHandler
|
||||
userService *service.UserService
|
||||
uploader *oss.UploaderManager
|
||||
}
|
||||
|
||||
func NewSunoHandler(app *core.AppServer, db *gorm.DB, userService *service.UserService, manager *oss.UploaderManager) *SunoHandler {
|
||||
return &SunoHandler{BaseHandler: handler.BaseHandler{App: app, DB: db}, userService: userService, uploader: manager}
|
||||
}
|
||||
|
||||
// RegisterRoutes 注册路由
|
||||
func (h *SunoHandler) RegisterRoutes() {
|
||||
group := h.App.Engine.Group("/api/admin/suno/")
|
||||
|
||||
// 需要管理员授权的接口
|
||||
group.Use(middleware.AdminAuthMiddleware(h.App.Config.AdminSession.SecretKey, h.App.Redis))
|
||||
{
|
||||
group.POST("list", h.SunoList)
|
||||
group.GET("remove", h.Remove)
|
||||
}
|
||||
}
|
||||
|
||||
type sunoQuery struct {
|
||||
Title string `json:"title"`
|
||||
Prompt string `json:"prompt"`
|
||||
CreatedAt []string `json:"created_at"`
|
||||
Page int `json:"page"`
|
||||
PageSize int `json:"page_size"`
|
||||
}
|
||||
|
||||
// SunoList Suno 任务列表
|
||||
func (h *SunoHandler) SunoList(c *gin.Context) {
|
||||
var data sunoQuery
|
||||
if err := c.ShouldBindJSON(&data); err != nil {
|
||||
resp.ERROR(c, types.InvalidArgs)
|
||||
return
|
||||
}
|
||||
|
||||
session := h.DB.Session(&gorm.Session{})
|
||||
if data.Title != "" {
|
||||
session = session.Where("title LIKE ?", "%"+data.Title+"%")
|
||||
}
|
||||
if data.Prompt != "" {
|
||||
// 同时查询 prompt 字段和 params JSON 字段中的 prompt
|
||||
session = session.Where("prompt LIKE ? OR JSON_EXTRACT(params, '$.prompt') LIKE ?", "%"+data.Prompt+"%", "%"+data.Prompt+"%")
|
||||
}
|
||||
if len(data.CreatedAt) == 2 {
|
||||
session = session.Where("created_at >= ? AND created_at <= ?", data.CreatedAt[0], data.CreatedAt[1])
|
||||
}
|
||||
var total int64
|
||||
session.Model(&model.SunoJob{}).Count(&total)
|
||||
var list []model.SunoJob
|
||||
var items = make([]vo.SunoJob, 0)
|
||||
offset := (data.Page - 1) * data.PageSize
|
||||
err := session.Order("id DESC").Offset(offset).Limit(data.PageSize).Find(&list).Error
|
||||
if err == nil {
|
||||
// 填充数据
|
||||
for _, item := range list {
|
||||
var job vo.SunoJob
|
||||
err = utils.CopyObject(item, &job)
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
job.CreatedAt = item.CreatedAt.Unix()
|
||||
items = append(items, job)
|
||||
}
|
||||
}
|
||||
|
||||
resp.SUCCESS(c, vo.NewPage(total, data.Page, data.PageSize, items))
|
||||
}
|
||||
|
||||
func (h *SunoHandler) Remove(c *gin.Context) {
|
||||
id := h.GetInt(c, "id", 0)
|
||||
|
||||
tx := h.DB.Begin()
|
||||
var job model.SunoJob
|
||||
if err := h.DB.Where("id", id).First(&job).Error; err != nil {
|
||||
resp.ERROR(c, "记录不存在")
|
||||
return
|
||||
}
|
||||
|
||||
// 删除任务
|
||||
tx.Delete(&job)
|
||||
md := "suno"
|
||||
power := job.Power
|
||||
userId := int(job.UserId)
|
||||
remark := fmt.Sprintf("SUNO 任务失败,退回算力。任务ID:%d,Err: %s", job.Id, job.ErrMsg)
|
||||
needRefund := job.Progress != 100
|
||||
fileURL := job.AudioURL
|
||||
|
||||
if needRefund {
|
||||
err := h.userService.IncreasePower(uint(userId), power, model.PowerLog{
|
||||
Type: types.PowerRefund,
|
||||
Model: md,
|
||||
Remark: remark,
|
||||
})
|
||||
if err != nil {
|
||||
tx.Rollback()
|
||||
resp.ERROR(c, err.Error())
|
||||
return
|
||||
}
|
||||
}
|
||||
tx.Commit()
|
||||
// remove file
|
||||
err := h.uploader.GetUploadHandler().Delete(fileURL)
|
||||
if err != nil {
|
||||
logger.Error("remove file failed: ", err)
|
||||
}
|
||||
|
||||
resp.SUCCESS(c)
|
||||
}
|
||||
@@ -8,6 +8,7 @@ package admin
|
||||
// * +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"fmt"
|
||||
"geekai/core"
|
||||
"geekai/core/middleware"
|
||||
@@ -17,11 +18,16 @@ import (
|
||||
"geekai/store/vo"
|
||||
"geekai/utils"
|
||||
"geekai/utils/resp"
|
||||
"net/http"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/go-redis/redis/v8"
|
||||
"github.com/golang-jwt/jwt/v5"
|
||||
"github.com/xuri/excelize/v2"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
@@ -47,9 +53,220 @@ func (h *UserHandler) RegisterRoutes() {
|
||||
group.GET("loginLog", h.LoginLog)
|
||||
group.GET("genLoginLink", h.GenLoginLink)
|
||||
group.POST("resetPass", h.ResetPass)
|
||||
group.GET("import/template", h.ImportTemplate)
|
||||
group.POST("import", h.ImportUsers)
|
||||
}
|
||||
}
|
||||
|
||||
// ImportTemplate 下载用户导入模板
|
||||
func (h *UserHandler) ImportTemplate(c *gin.Context) {
|
||||
f := excelize.NewFile()
|
||||
sheetName := "Sheet1"
|
||||
// 表头
|
||||
headers := []string{"用户名", "密码", "手机", "邮箱", "剩余算力", "启用状态"}
|
||||
for i, title := range headers {
|
||||
cell, _ := excelize.CoordinatesToCellName(i+1, 1)
|
||||
_ = f.SetCellValue(sheetName, cell, title)
|
||||
}
|
||||
// 示例数据
|
||||
sample := []interface{}{"user001", "Passw0rd!", "13800000000", "user001@example.com", 100, 1}
|
||||
for i, v := range sample {
|
||||
cell, _ := excelize.CoordinatesToCellName(i+1, 2)
|
||||
_ = f.SetCellValue(sheetName, cell, v)
|
||||
}
|
||||
|
||||
buf, err := f.WriteToBuffer()
|
||||
if err != nil {
|
||||
logger.Error("failed to generate user import template: ", err)
|
||||
resp.ERROR(c, "生成模板失败")
|
||||
return
|
||||
}
|
||||
|
||||
c.Header("Content-Type", "application/vnd.openxmlformats-officedocument.spreadsheetml.sheet")
|
||||
c.Header("Content-Disposition", "attachment; filename=\"user_import_template.xlsx\"")
|
||||
c.Data(http.StatusOK, "application/vnd.openxmlformats-officedocument.spreadsheetml.sheet", buf.Bytes())
|
||||
}
|
||||
|
||||
// ImportUsers 批量导入用户
|
||||
func (h *UserHandler) ImportUsers(c *gin.Context) {
|
||||
fileHeader, err := c.FormFile("file")
|
||||
if err != nil {
|
||||
resp.ERROR(c, "文件上传失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
ext := strings.ToLower(filepath.Ext(fileHeader.Filename))
|
||||
if ext != ".xlsx" {
|
||||
resp.ERROR(c, "只支持 .xlsx 格式的 Excel 文件")
|
||||
return
|
||||
}
|
||||
|
||||
src, err := fileHeader.Open()
|
||||
if err != nil {
|
||||
resp.ERROR(c, "无法读取上传文件: "+err.Error())
|
||||
return
|
||||
}
|
||||
defer src.Close()
|
||||
|
||||
// 读取到内存,避免多次读取问题
|
||||
var buf bytes.Buffer
|
||||
if _, err = buf.ReadFrom(src); err != nil {
|
||||
resp.ERROR(c, "读取文件内容失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
excel, err := excelize.OpenReader(bytes.NewReader(buf.Bytes()))
|
||||
if err != nil {
|
||||
resp.ERROR(c, "解析 Excel 失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
defer func() {
|
||||
_ = excel.Close()
|
||||
}()
|
||||
|
||||
rows, err := excel.GetRows("Sheet1")
|
||||
if err != nil {
|
||||
resp.ERROR(c, "读取工作表失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
if len(rows) < 2 {
|
||||
resp.ERROR(c, "Excel 中没有可导入的数据")
|
||||
return
|
||||
}
|
||||
|
||||
type rowError struct {
|
||||
Row int `json:"row"`
|
||||
Error string `json:"error"`
|
||||
}
|
||||
|
||||
var (
|
||||
successCount int
|
||||
failedCount int
|
||||
errorsList []rowError
|
||||
)
|
||||
|
||||
usernameSet := make(map[string]struct{})
|
||||
|
||||
// 从第二行开始读取
|
||||
for index, row := range rows[1:] {
|
||||
line := index + 2 // Excel 行号
|
||||
if len(row) == 0 {
|
||||
continue
|
||||
}
|
||||
|
||||
get := func(i int) string {
|
||||
if i < len(row) {
|
||||
return strings.TrimSpace(row[i])
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
username := get(0)
|
||||
password := get(1)
|
||||
mobile := get(2)
|
||||
email := get(3)
|
||||
powerStr := get(4)
|
||||
statusStr := get(5)
|
||||
|
||||
// 基础校验
|
||||
if username == "" {
|
||||
failedCount++
|
||||
errorsList = append(errorsList, rowError{Row: line, Error: "用户名不能为空"})
|
||||
continue
|
||||
}
|
||||
if _, ok := usernameSet[username]; ok {
|
||||
failedCount++
|
||||
errorsList = append(errorsList, rowError{Row: line, Error: "同一文件中用户名重复"})
|
||||
continue
|
||||
}
|
||||
usernameSet[username] = struct{}{}
|
||||
|
||||
if len(password) < 8 || len(password) > 16 {
|
||||
failedCount++
|
||||
errorsList = append(errorsList, rowError{Row: line, Error: "密码必须为 8-16 位"})
|
||||
continue
|
||||
}
|
||||
|
||||
if mobile != "" && len(mobile) != 11 {
|
||||
failedCount++
|
||||
errorsList = append(errorsList, rowError{Row: line, Error: "手机号必须为 11 位"})
|
||||
continue
|
||||
}
|
||||
|
||||
// 解析算力
|
||||
power := 0
|
||||
if powerStr != "" {
|
||||
p, err := strconv.Atoi(powerStr)
|
||||
if err != nil {
|
||||
failedCount++
|
||||
errorsList = append(errorsList, rowError{Row: line, Error: "剩余算力必须为数字"})
|
||||
continue
|
||||
}
|
||||
if p < 0 {
|
||||
failedCount++
|
||||
errorsList = append(errorsList, rowError{Row: line, Error: "剩余算力不能为负数"})
|
||||
continue
|
||||
}
|
||||
power = p
|
||||
}
|
||||
|
||||
// 解析启用状态
|
||||
status := true
|
||||
if statusStr != "" {
|
||||
switch strings.TrimSpace(statusStr) {
|
||||
case "0", "否", "false", "停用":
|
||||
status = false
|
||||
case "1", "是", "true", "启用":
|
||||
status = true
|
||||
default:
|
||||
failedCount++
|
||||
errorsList = append(errorsList, rowError{Row: line, Error: "启用状态只支持 1/是 或 0/否"})
|
||||
continue
|
||||
}
|
||||
}
|
||||
|
||||
// 检查用户名是否已存在
|
||||
var exist model.User
|
||||
if err = h.DB.Where("username = ?", username).First(&exist).Error; err == nil && exist.Id > 0 {
|
||||
failedCount++
|
||||
errorsList = append(errorsList, rowError{Row: line, Error: "用户名已存在"})
|
||||
continue
|
||||
}
|
||||
|
||||
salt := utils.RandString(8)
|
||||
u := model.User{
|
||||
Username: username,
|
||||
Password: utils.GenPassword(password, salt),
|
||||
Mobile: mobile,
|
||||
Email: email,
|
||||
Avatar: "/images/avatar/user.png",
|
||||
Salt: salt,
|
||||
Power: power,
|
||||
Status: status,
|
||||
ChatRoles: utils.JsonEncode([]string{}),
|
||||
ChatConfig: "{}",
|
||||
ChatModels: utils.JsonEncode([]int{}),
|
||||
ExpiredTime: 0, // 长期有效
|
||||
Vip: false,
|
||||
}
|
||||
u.Nickname = fmt.Sprintf("用户@%d", utils.RandomNumber(6))
|
||||
|
||||
if err = h.DB.Create(&u).Error; err != nil {
|
||||
failedCount++
|
||||
errorsList = append(errorsList, rowError{Row: line, Error: "写入数据库失败: " + err.Error()})
|
||||
continue
|
||||
}
|
||||
|
||||
successCount++
|
||||
}
|
||||
|
||||
resp.SUCCESS(c, gin.H{
|
||||
"success": successCount,
|
||||
"failed": failedCount,
|
||||
"errors": errorsList,
|
||||
})
|
||||
}
|
||||
|
||||
// List 用户列表
|
||||
func (h *UserHandler) List(c *gin.Context) {
|
||||
page := h.GetInt(c, "page", 1)
|
||||
@@ -96,17 +313,16 @@ func (h *UserHandler) List(c *gin.Context) {
|
||||
|
||||
func (h *UserHandler) Save(c *gin.Context) {
|
||||
var data struct {
|
||||
Id uint `json:"id"`
|
||||
Password string `json:"password"`
|
||||
Username string `json:"username"`
|
||||
Mobile string `json:"mobile"`
|
||||
Email string `json:"email"`
|
||||
ChatRoles []string `json:"chat_roles"`
|
||||
ChatModels []int `json:"chat_models"`
|
||||
ExpiredTime string `json:"expired_time"`
|
||||
Status bool `json:"status"`
|
||||
Vip bool `json:"vip"`
|
||||
Power int `json:"power"`
|
||||
Id uint `json:"id"`
|
||||
Password string `json:"password"`
|
||||
Username string `json:"username"`
|
||||
Mobile string `json:"mobile"`
|
||||
Email string `json:"email"`
|
||||
ChatModels []int `json:"chat_models"`
|
||||
ExpiredTime string `json:"expired_time"`
|
||||
Status bool `json:"status"`
|
||||
Vip bool `json:"vip"`
|
||||
Power int `json:"power"`
|
||||
}
|
||||
if err := c.ShouldBindJSON(&data); err != nil {
|
||||
resp.ERROR(c, types.InvalidArgs)
|
||||
@@ -128,11 +344,10 @@ func (h *UserHandler) Save(c *gin.Context) {
|
||||
user.Status = data.Status
|
||||
user.Vip = data.Vip
|
||||
user.Power = data.Power
|
||||
user.ChatRoles = utils.JsonEncode(data.ChatRoles)
|
||||
user.ChatModels = utils.JsonEncode(data.ChatModels)
|
||||
user.ExpiredTime = utils.Str2stamp(data.ExpiredTime)
|
||||
|
||||
res = h.DB.Select("username", "mobile", "email", "status", "vip", "power", "chat_roles_json", "chat_models_json", "expired_time").Updates(&user)
|
||||
res = h.DB.Select("username", "mobile", "email", "status", "vip", "power", "chat_models_json", "expired_time").Updates(&user)
|
||||
|
||||
if res.Error != nil {
|
||||
logger.Error("error with update database:", res.Error)
|
||||
@@ -184,7 +399,6 @@ func (h *UserHandler) Save(c *gin.Context) {
|
||||
Salt: salt,
|
||||
Power: data.Power,
|
||||
Status: true,
|
||||
ChatRoles: utils.JsonEncode(data.ChatRoles),
|
||||
ChatConfig: "{}",
|
||||
ChatModels: utils.JsonEncode(data.ChatModels),
|
||||
ExpiredTime: utils.Str2stamp(data.ExpiredTime),
|
||||
@@ -278,10 +492,7 @@ func (h *UserHandler) Remove(c *gin.Context) {
|
||||
if err = tx.Where("user_id = ?", id).Delete(&model.MidJourneyJob{}).Error; err != nil {
|
||||
break
|
||||
}
|
||||
if err = tx.Where("user_id = ?", id).Delete(&model.SdJob{}).Error; err != nil {
|
||||
break
|
||||
}
|
||||
if err = tx.Where("user_id = ?", id).Delete(&model.DallJob{}).Error; err != nil {
|
||||
if err = tx.Where("user_id = ?", id).Delete(&model.ImageJob{}).Error; err != nil {
|
||||
break
|
||||
}
|
||||
if err = tx.Where("user_id = ?", id).Delete(&model.SunoJob{}).Error; err != nil {
|
||||
|
||||
@@ -0,0 +1,269 @@
|
||||
package admin
|
||||
|
||||
// * +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
|
||||
// * Copyright 2023 The Geek-AI Authors. All rights reserved.
|
||||
// * Use of this source code is governed by a Apache-2.0 license
|
||||
// * that can be found in the LICENSE file.
|
||||
// * @Author yangjian102621@163.com
|
||||
// * +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"geekai/core"
|
||||
"geekai/core/middleware"
|
||||
"geekai/core/types"
|
||||
"geekai/handler"
|
||||
"geekai/service"
|
||||
"geekai/service/oss"
|
||||
"geekai/store/model"
|
||||
"geekai/store/vo"
|
||||
"geekai/utils"
|
||||
"geekai/utils/resp"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// VideoHandler 管理后台视频生成处理器
|
||||
type VideoHandler struct {
|
||||
handler.BaseHandler
|
||||
userService *service.UserService
|
||||
uploader *oss.UploaderManager
|
||||
}
|
||||
|
||||
// NewVideoHandler 创建管理后台视频生成处理器
|
||||
func NewVideoHandler(app *core.AppServer, db *gorm.DB, userService *service.UserService, manager *oss.UploaderManager) *VideoHandler {
|
||||
return &VideoHandler{
|
||||
BaseHandler: handler.BaseHandler{App: app, DB: db},
|
||||
userService: userService,
|
||||
uploader: manager,
|
||||
}
|
||||
}
|
||||
|
||||
// RegisterRoutes 注册视频生成管理后台路由
|
||||
func (h *VideoHandler) RegisterRoutes() {
|
||||
rg := h.App.Engine.Group("/api/admin/video/")
|
||||
rg.Use(middleware.AdminAuthMiddleware(h.App.Config.AdminSession.SecretKey, h.App.Redis))
|
||||
{
|
||||
rg.GET("config", h.GetConfig)
|
||||
rg.POST("config/update", h.UpdateConfig)
|
||||
rg.POST("list", h.Videos)
|
||||
rg.GET("remove", h.Remove)
|
||||
}
|
||||
}
|
||||
|
||||
// GetConfig 获取视频生成配置
|
||||
func (h *VideoHandler) GetConfig(c *gin.Context) {
|
||||
var config model.Config
|
||||
err := h.DB.Where("name", types.ConfigKeyVideo).First(&config).Error
|
||||
if err != nil {
|
||||
if err == gorm.ErrRecordNotFound {
|
||||
// 返回空配置
|
||||
resp.SUCCESS(c, types.VideoConfig{
|
||||
ApiURL: "",
|
||||
ApiKey: "",
|
||||
VideoPowers: make(map[string]types.VideoModelPower),
|
||||
})
|
||||
return
|
||||
}
|
||||
resp.ERROR(c, "获取配置失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
var videoConfig types.VideoConfig
|
||||
err = utils.JsonDecode(config.Value, &videoConfig)
|
||||
if err != nil {
|
||||
resp.ERROR(c, "解析配置失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
resp.SUCCESS(c, videoConfig)
|
||||
}
|
||||
|
||||
// UpdateConfig 更新视频生成配置
|
||||
func (h *VideoHandler) UpdateConfig(c *gin.Context) {
|
||||
var req types.VideoConfig
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
resp.ERROR(c, "参数错误")
|
||||
return
|
||||
}
|
||||
|
||||
// 验证必填字段
|
||||
if req.ApiURL == "" {
|
||||
resp.ERROR(c, "API地址不能为空")
|
||||
return
|
||||
}
|
||||
if req.ApiKey == "" {
|
||||
resp.ERROR(c, "API密钥不能为空")
|
||||
return
|
||||
}
|
||||
|
||||
// 验证算力配置
|
||||
if len(req.VideoPowers) == 0 {
|
||||
resp.ERROR(c, "请至少配置一个模型的算力")
|
||||
return
|
||||
}
|
||||
|
||||
// 新的价格配置方式直接使用 power_config 中的 key(如 "fixed"、"5_720P" 等)
|
||||
// 不再区分固定收费和按秒收费,所有价格配置都在 power_config 中
|
||||
for key, modelPower := range req.VideoPowers {
|
||||
// 验证 provider
|
||||
if modelPower.Provider == "" {
|
||||
resp.ERROR(c, fmt.Sprintf("模型 %s 的 provider 不能为空", key))
|
||||
return
|
||||
}
|
||||
|
||||
// 验证 model
|
||||
if modelPower.Model == "" {
|
||||
resp.ERROR(c, fmt.Sprintf("模型 %s 的 model 不能为空", key))
|
||||
return
|
||||
}
|
||||
|
||||
// 验证 power_config
|
||||
if len(modelPower.PowerConfig) == 0 {
|
||||
resp.ERROR(c, fmt.Sprintf("模型 %s 的 power_config 不能为空", key))
|
||||
return
|
||||
}
|
||||
|
||||
// 验证 power_config 中的值必须大于0
|
||||
for configKey, configValue := range modelPower.PowerConfig {
|
||||
if configValue <= 0 {
|
||||
resp.ERROR(c, fmt.Sprintf("模型 %s 的 power_config.%s 必须大于0", key, configKey))
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
// 保存配置
|
||||
tx := h.DB.Begin()
|
||||
value := utils.JsonEncode(&req)
|
||||
var exist model.Config
|
||||
tx.Where("name", types.ConfigKeyVideo).First(&exist)
|
||||
|
||||
if exist.Id > 0 {
|
||||
exist.Value = value
|
||||
err := tx.Updates(&exist).Error
|
||||
if err != nil {
|
||||
resp.ERROR(c, "更新配置失败: "+err.Error())
|
||||
tx.Rollback()
|
||||
return
|
||||
}
|
||||
} else {
|
||||
exist.Name = types.ConfigKeyVideo
|
||||
exist.Value = value
|
||||
err := tx.Create(&exist).Error
|
||||
if err != nil {
|
||||
resp.ERROR(c, "创建配置失败: "+err.Error())
|
||||
tx.Rollback()
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
tx.Commit()
|
||||
resp.SUCCESS(c, gin.H{"message": "配置更新成功"})
|
||||
}
|
||||
|
||||
type videoQuery struct {
|
||||
Type string `json:"type"` // 任务类型 luma, keling
|
||||
Status string `json:"status"` // 任务状态 pending, in_progress, downloading, success, failed
|
||||
Prompt string `json:"prompt"`
|
||||
CreatedAt []string `json:"created_at"`
|
||||
Page int `json:"page"`
|
||||
PageSize int `json:"page_size"`
|
||||
}
|
||||
|
||||
// Videos 视频任务列表
|
||||
func (h *VideoHandler) Videos(c *gin.Context) {
|
||||
var data videoQuery
|
||||
if err := c.ShouldBindJSON(&data); err != nil {
|
||||
resp.ERROR(c, types.InvalidArgs)
|
||||
return
|
||||
}
|
||||
|
||||
session := h.DB.Session(&gorm.Session{})
|
||||
if data.Type != "" {
|
||||
session = session.Where("type", data.Type)
|
||||
}
|
||||
if data.Status != "" {
|
||||
session = session.Where("status", data.Status)
|
||||
}
|
||||
if data.Prompt != "" {
|
||||
session = session.Where("prompt LIKE ?", "%"+data.Prompt+"%")
|
||||
}
|
||||
if len(data.CreatedAt) == 2 {
|
||||
session = session.Where("created_at >= ? AND created_at <= ?", data.CreatedAt[0], data.CreatedAt[1])
|
||||
}
|
||||
var total int64
|
||||
session.Model(&model.VideoJob{}).Count(&total)
|
||||
var list []model.VideoJob
|
||||
var items = make([]vo.VideoJob, 0)
|
||||
offset := (data.Page - 1) * data.PageSize
|
||||
err := session.Order("id DESC").Offset(offset).Limit(data.PageSize).Find(&list).Error
|
||||
if err == nil {
|
||||
// 填充数据
|
||||
for _, item := range list {
|
||||
var job vo.VideoJob
|
||||
err = utils.CopyObject(item, &job)
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
job.CreatedAt = item.CreatedAt.Unix()
|
||||
items = append(items, job)
|
||||
}
|
||||
}
|
||||
|
||||
resp.SUCCESS(c, vo.NewPage(total, data.Page, data.PageSize, items))
|
||||
}
|
||||
|
||||
func (h *VideoHandler) Remove(c *gin.Context) {
|
||||
id := h.GetInt(c, "id", 0)
|
||||
tab := c.Query("tab")
|
||||
|
||||
tx := h.DB.Begin()
|
||||
var md, remark, fileURL string
|
||||
var power, userId int
|
||||
var needRefund bool
|
||||
|
||||
switch tab {
|
||||
case "luma", "keling":
|
||||
var job model.VideoJob
|
||||
if res := h.DB.Where("id", id).First(&job); res.Error != nil {
|
||||
resp.ERROR(c, "记录不存在")
|
||||
return
|
||||
}
|
||||
|
||||
// 删除任务
|
||||
tx.Delete(&job)
|
||||
md = job.Type
|
||||
power = job.Power
|
||||
userId = int(job.UserId)
|
||||
remark = fmt.Sprintf("视频任务失败,退回算力。任务ID:%d,Err: %s", job.Id, job.ErrMsg)
|
||||
needRefund = job.Status != types.VideoStatusSuccess
|
||||
fileURL = job.VideoURL
|
||||
default:
|
||||
resp.ERROR(c, types.InvalidArgs)
|
||||
return
|
||||
}
|
||||
|
||||
if needRefund {
|
||||
err := h.userService.IncreasePower(uint(userId), power, model.PowerLog{
|
||||
Type: types.PowerRefund,
|
||||
Model: md,
|
||||
Remark: remark,
|
||||
})
|
||||
if err != nil {
|
||||
tx.Rollback()
|
||||
resp.ERROR(c, err.Error())
|
||||
return
|
||||
}
|
||||
}
|
||||
tx.Commit()
|
||||
// remove file
|
||||
err := h.uploader.GetUploadHandler().Delete(fileURL)
|
||||
if err != nil {
|
||||
logger.Error("remove file failed: ", err)
|
||||
}
|
||||
|
||||
resp.SUCCESS(c)
|
||||
}
|
||||
@@ -12,7 +12,7 @@ import (
|
||||
"fmt"
|
||||
"geekai/core"
|
||||
"geekai/core/types"
|
||||
logger2 "geekai/logger"
|
||||
"geekai/log"
|
||||
"geekai/store/model"
|
||||
"geekai/utils"
|
||||
"strings"
|
||||
@@ -22,7 +22,7 @@ import (
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
var logger = logger2.GetLogger()
|
||||
var logger = log.GetLogger()
|
||||
|
||||
type BaseHandler struct {
|
||||
App *core.AppServer
|
||||
|
||||
+166
-31
@@ -37,7 +37,11 @@ func (h *ChatAppHandler) RegisterRoutes() {
|
||||
group.Use(middleware.UserAuthMiddleware(h.App.Config.Session.SecretKey, h.App.Redis))
|
||||
{
|
||||
group.GET("list/user", h.ListByUser)
|
||||
group.POST("create", h.Create)
|
||||
group.POST("copy", h.Copy)
|
||||
group.POST("update", h.UpdateApp)
|
||||
group.POST("workspace", h.UpdateWorkArea)
|
||||
group.POST("remove", h.Remove)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -45,7 +49,7 @@ func (h *ChatAppHandler) RegisterRoutes() {
|
||||
func (h *ChatAppHandler) List(c *gin.Context) {
|
||||
tid := h.GetInt(c, "tid", 0)
|
||||
var roles []model.ChatApp
|
||||
session := h.DB.Where("enable", true)
|
||||
session := h.DB.Where("enable = ? AND user_id = 0", true)
|
||||
if tid > 0 {
|
||||
session = session.Where("tid", tid)
|
||||
}
|
||||
@@ -61,6 +65,9 @@ func (h *ChatAppHandler) List(c *gin.Context) {
|
||||
err := utils.CopyObject(r, &v)
|
||||
if err == nil {
|
||||
v.Id = r.Id
|
||||
if r.UserId == 0 {
|
||||
v.SystemPrompt = ""
|
||||
}
|
||||
roleVos = append(roleVos, v)
|
||||
}
|
||||
}
|
||||
@@ -72,23 +79,11 @@ func (h *ChatAppHandler) ListByUser(c *gin.Context) {
|
||||
id := h.GetInt(c, "id", 0)
|
||||
userId := h.GetLoginUserId(c)
|
||||
var roles []model.ChatApp
|
||||
session := h.DB.Where("enable", true)
|
||||
// 如果用户没登录,则获取所有角色
|
||||
session := h.DB.Where("enable = ?", true)
|
||||
if userId > 0 {
|
||||
var user model.User
|
||||
h.DB.First(&user, userId)
|
||||
var roleKeys []string
|
||||
if user.ChatRoles != "" {
|
||||
err := utils.JsonDecode(user.ChatRoles, &roleKeys)
|
||||
if err != nil {
|
||||
resp.ERROR(c, "角色解析失败!")
|
||||
return
|
||||
}
|
||||
}
|
||||
// 保证用户至少有一个角色可用
|
||||
if len(roleKeys) > 0 {
|
||||
session = session.Where("marker IN ?", roleKeys)
|
||||
}
|
||||
session = session.Where("(user_id = 0 OR user_id = ?)", userId)
|
||||
} else {
|
||||
session = session.Where("user_id = 0")
|
||||
}
|
||||
|
||||
if id > 0 {
|
||||
@@ -106,33 +101,173 @@ func (h *ChatAppHandler) ListByUser(c *gin.Context) {
|
||||
err := utils.CopyObject(r, &v)
|
||||
if err == nil {
|
||||
v.Id = r.Id
|
||||
if r.UserId == 0 {
|
||||
v.SystemPrompt = ""
|
||||
}
|
||||
roleVos = append(roleVos, v)
|
||||
}
|
||||
}
|
||||
resp.SUCCESS(c, roleVos)
|
||||
}
|
||||
|
||||
// UpdateApp 更新用户聊天应用
|
||||
func (h *ChatAppHandler) UpdateApp(c *gin.Context) {
|
||||
user, err := h.GetLoginUser(c)
|
||||
if err != nil {
|
||||
// Create 用户创建智能体
|
||||
func (h *ChatAppHandler) Create(c *gin.Context) {
|
||||
userId := h.GetLoginUserId(c)
|
||||
if userId == 0 {
|
||||
resp.NotAuth(c)
|
||||
return
|
||||
}
|
||||
|
||||
var data struct {
|
||||
Keys []string `json:"keys"`
|
||||
}
|
||||
if err = c.ShouldBindJSON(&data); err != nil {
|
||||
var data vo.ChatApp
|
||||
if err := c.ShouldBindJSON(&data); err != nil {
|
||||
resp.ERROR(c, types.InvalidArgs)
|
||||
return
|
||||
}
|
||||
|
||||
err = h.DB.Model(&model.User{}).Where("id = ?", user.Id).UpdateColumn("chat_roles_json", utils.JsonEncode(data.Keys)).Error
|
||||
if err != nil {
|
||||
role := model.ChatApp{
|
||||
Name: data.Name,
|
||||
Tid: data.Tid,
|
||||
UserId: userId,
|
||||
SystemPrompt: data.SystemPrompt,
|
||||
HelloMsg: data.HelloMsg,
|
||||
Icon: data.Icon,
|
||||
Enable: true,
|
||||
SortNum: int(data.SortNum),
|
||||
ModelId: data.ModelId,
|
||||
}
|
||||
if role.Icon == "" {
|
||||
role.Icon = "/images/avatar/gpt.png"
|
||||
}
|
||||
if err := h.DB.Create(&role).Error; err != nil {
|
||||
resp.ERROR(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
resp.SUCCESS(c)
|
||||
data.Id = role.Id
|
||||
data.UserId = role.UserId
|
||||
resp.SUCCESS(c, data)
|
||||
}
|
||||
|
||||
// Copy 用户复制智能体(复制为当前用户名下)
|
||||
func (h *ChatAppHandler) Copy(c *gin.Context) {
|
||||
userId := h.GetLoginUserId(c)
|
||||
if userId == 0 {
|
||||
resp.NotAuth(c)
|
||||
return
|
||||
}
|
||||
var body struct {
|
||||
SourceId uint `json:"source_id"`
|
||||
}
|
||||
if err := c.ShouldBindJSON(&body); err != nil || body.SourceId == 0 {
|
||||
resp.ERROR(c, types.InvalidArgs)
|
||||
return
|
||||
}
|
||||
var src model.ChatApp
|
||||
if err := h.DB.First(&src, body.SourceId).Error; err != nil {
|
||||
resp.ERROR(c, "智能体不存在")
|
||||
return
|
||||
}
|
||||
role := model.ChatApp{
|
||||
Name: src.Name,
|
||||
Tid: src.Tid,
|
||||
UserId: userId,
|
||||
SystemPrompt: src.SystemPrompt,
|
||||
HelloMsg: src.HelloMsg,
|
||||
Icon: src.Icon,
|
||||
Enable: true,
|
||||
SortNum: src.SortNum,
|
||||
ModelId: src.ModelId,
|
||||
}
|
||||
if err := h.DB.Create(&role).Error; err != nil {
|
||||
resp.ERROR(c, err.Error())
|
||||
return
|
||||
}
|
||||
resp.SUCCESS(c, gin.H{"id": role.Id})
|
||||
}
|
||||
|
||||
// UpdateApp 更新用户聊天应用(仅允许更新自己创建的)
|
||||
func (h *ChatAppHandler) UpdateApp(c *gin.Context) {
|
||||
userId := h.GetLoginUserId(c)
|
||||
if userId == 0 {
|
||||
resp.NotAuth(c)
|
||||
return
|
||||
}
|
||||
var data vo.ChatApp
|
||||
if err := c.ShouldBindJSON(&data); err != nil || data.Id == 0 {
|
||||
resp.ERROR(c, types.InvalidArgs)
|
||||
return
|
||||
}
|
||||
var role model.ChatApp
|
||||
if err := h.DB.First(&role, data.Id).Error; err != nil {
|
||||
resp.ERROR(c, "智能体不存在")
|
||||
return
|
||||
}
|
||||
if role.UserId != userId {
|
||||
resp.ERROR(c, "无权限修改该智能体")
|
||||
return
|
||||
}
|
||||
updates := map[string]interface{}{
|
||||
"name": data.Name,
|
||||
"hello_msg": data.HelloMsg,
|
||||
"icon": data.Icon,
|
||||
"model_id": data.ModelId,
|
||||
"system_prompt": data.SystemPrompt,
|
||||
}
|
||||
if err := h.DB.Model(&role).Updates(updates).Error; err != nil {
|
||||
resp.ERROR(c, err.Error())
|
||||
return
|
||||
}
|
||||
resp.SUCCESS(c, nil)
|
||||
}
|
||||
|
||||
// UpdateWorkArea 更新用户工作区应用列表(存为应用 id 数组)
|
||||
func (h *ChatAppHandler) UpdateWorkArea(c *gin.Context) {
|
||||
userId := h.GetLoginUserId(c)
|
||||
if userId == 0 {
|
||||
resp.NotAuth(c)
|
||||
return
|
||||
}
|
||||
var body struct {
|
||||
Ids []uint `json:"ids"`
|
||||
}
|
||||
if err := c.ShouldBindJSON(&body); err != nil {
|
||||
resp.ERROR(c, types.InvalidArgs)
|
||||
return
|
||||
}
|
||||
if err := h.DB.Model(&model.User{}).Where("id = ?", userId).Update("chat_roles_json", utils.JsonEncode(body.Ids)).Error; err != nil {
|
||||
resp.ERROR(c, err.Error())
|
||||
return
|
||||
}
|
||||
resp.SUCCESS(c, nil)
|
||||
}
|
||||
|
||||
// Remove 删除用户智能体(仅允许删除自己创建的)
|
||||
func (h *ChatAppHandler) Remove(c *gin.Context) {
|
||||
userId := h.GetLoginUserId(c)
|
||||
if userId == 0 {
|
||||
resp.NotAuth(c)
|
||||
return
|
||||
}
|
||||
var body struct {
|
||||
Id uint `json:"id"`
|
||||
}
|
||||
_ = c.ShouldBindJSON(&body)
|
||||
if body.Id == 0 {
|
||||
body.Id = uint(h.GetInt(c, "id", 0))
|
||||
}
|
||||
if body.Id == 0 {
|
||||
resp.ERROR(c, types.InvalidArgs)
|
||||
return
|
||||
}
|
||||
var role model.ChatApp
|
||||
if err := h.DB.First(&role, body.Id).Error; err != nil {
|
||||
resp.ERROR(c, "智能体不存在")
|
||||
return
|
||||
}
|
||||
if role.UserId != userId {
|
||||
resp.ERROR(c, "无权限删除该智能体")
|
||||
return
|
||||
}
|
||||
if err := h.DB.Delete(&role).Error; err != nil {
|
||||
resp.ERROR(c, err.Error())
|
||||
return
|
||||
}
|
||||
resp.SUCCESS(c, nil)
|
||||
}
|
||||
|
||||
@@ -95,20 +95,21 @@ func NewChatHandler(app *core.AppServer,
|
||||
// RegisterRoutes 注册路由
|
||||
func (h *ChatHandler) RegisterRoutes() {
|
||||
group := h.App.Engine.Group("/api/chat/")
|
||||
group.GET("detail", h.Detail)
|
||||
group.GET("history", h.History)
|
||||
// 其他接口需要用户授权
|
||||
group.Use(middleware.UserAuthMiddleware(h.App.Config.Session.SecretKey, h.App.Redis))
|
||||
{
|
||||
group.Any("message", h.Chat)
|
||||
group.GET("list", h.List)
|
||||
group.GET("detail", h.Detail)
|
||||
group.POST("update", h.Update)
|
||||
group.GET("remove", h.Remove)
|
||||
group.GET("history", h.History)
|
||||
group.GET("clear", h.Clear)
|
||||
group.POST("tokens", h.Tokens)
|
||||
group.GET("stop", h.StopGenerate)
|
||||
group.POST("tts", h.TextToSpeech)
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
// Chat 处理聊天请求
|
||||
@@ -272,7 +273,7 @@ func (h *ChatHandler) sendMessage(ctx context.Context, input ChatInput, c *gin.C
|
||||
chatCtx := make([]any, 0)
|
||||
messages := make([]any, 0)
|
||||
if h.App.SysConfig.Base.EnableContext {
|
||||
_ = utils.JsonDecode(input.ChatRole.Context, &messages)
|
||||
_ = utils.JsonDecode(input.ChatRole.SystemPrompt, &messages)
|
||||
if h.App.SysConfig.Base.ContextDeep > 0 {
|
||||
var historyMessages []model.ChatMessage
|
||||
dbSession := h.DB.Session(&gorm.Session{}).Where("chat_id", input.ChatId)
|
||||
@@ -668,7 +669,10 @@ func (h *ChatHandler) saveChatHistory(
|
||||
files := make([]vo.File, 0)
|
||||
if strings.HasPrefix(req.Model, "sora") {
|
||||
video, err := h.soraService.DownloadVideoURL(message.Content)
|
||||
if err == nil {
|
||||
if err != nil {
|
||||
logger.Error("failed to download video: ", err)
|
||||
pushMessage(c, ChatEventError, "视频下载失败:"+err.Error())
|
||||
} else {
|
||||
files = append(files, *video)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -8,7 +8,10 @@ package handler
|
||||
// * +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"geekai/core"
|
||||
"geekai/core/types"
|
||||
"geekai/service/oss"
|
||||
"geekai/store/model"
|
||||
"geekai/utils"
|
||||
"geekai/utils/resp"
|
||||
@@ -19,10 +22,16 @@ import (
|
||||
|
||||
type ConfigHandler struct {
|
||||
BaseHandler
|
||||
uploaderManager *oss.UploaderManager
|
||||
sysConfig *types.SystemConfig
|
||||
}
|
||||
|
||||
func NewConfigHandler(app *core.AppServer, db *gorm.DB) *ConfigHandler {
|
||||
return &ConfigHandler{BaseHandler: BaseHandler{App: app, DB: db}}
|
||||
func NewConfigHandler(app *core.AppServer, db *gorm.DB, uploaderManager *oss.UploaderManager, sysConfig *types.SystemConfig) *ConfigHandler {
|
||||
return &ConfigHandler{
|
||||
BaseHandler: BaseHandler{App: app, DB: db},
|
||||
uploaderManager: uploaderManager,
|
||||
sysConfig: sysConfig,
|
||||
}
|
||||
}
|
||||
|
||||
// RegisterRoutes 注册路由
|
||||
@@ -31,24 +40,44 @@ func (h *ConfigHandler) RegisterRoutes() {
|
||||
|
||||
// 无需授权的接口
|
||||
group.GET("get", h.Get)
|
||||
group.GET("oss/thumb", h.GetOssThumbTemplate)
|
||||
}
|
||||
|
||||
// Get 获取指定的系统配置
|
||||
func (h *ConfigHandler) Get(c *gin.Context) {
|
||||
key := c.Query("key")
|
||||
var config model.Config
|
||||
res := h.DB.Where("name", key).First(&config)
|
||||
if res.Error != nil {
|
||||
resp.ERROR(c, res.Error.Error())
|
||||
err := h.DB.Where("name", key).First(&config).Error
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
resp.SUCCESS(c, nil)
|
||||
return
|
||||
}
|
||||
|
||||
var value map[string]any
|
||||
err := utils.JsonDecode(config.Value, &value)
|
||||
if err != nil {
|
||||
resp.ERROR(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
var value map[string]any
|
||||
err = utils.JsonDecode(config.Value, &value)
|
||||
if err != nil {
|
||||
resp.ERROR(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
if key == types.ConfigKeyWxGzh {
|
||||
delete(value, "secret")
|
||||
delete(value, "token")
|
||||
delete(value, "encoding_aes_key")
|
||||
}
|
||||
resp.SUCCESS(c, value)
|
||||
}
|
||||
|
||||
// GetOssThumbTemplate 获取当前存储引擎的缩略图模板
|
||||
func (h *ConfigHandler) GetOssThumbTemplate(c *gin.Context) {
|
||||
template := h.uploaderManager.GetThumbTemplate()
|
||||
resp.SUCCESS(c, gin.H{
|
||||
"template": template,
|
||||
"active": h.sysConfig.OSS.Active,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -13,7 +13,7 @@ import (
|
||||
"geekai/core"
|
||||
"geekai/core/types"
|
||||
"geekai/service"
|
||||
"geekai/service/dalle"
|
||||
"geekai/service/image"
|
||||
"geekai/service/oss"
|
||||
"geekai/store/model"
|
||||
"geekai/store/vo"
|
||||
@@ -31,7 +31,7 @@ import (
|
||||
type FunctionHandler struct {
|
||||
BaseHandler
|
||||
uploadManager *oss.UploaderManager
|
||||
dallService *dalle.Service
|
||||
imageService *image.Service
|
||||
userService *service.UserService
|
||||
}
|
||||
|
||||
@@ -40,7 +40,7 @@ func NewFunctionHandler(
|
||||
db *gorm.DB,
|
||||
config *types.AppConfig,
|
||||
manager *oss.UploaderManager,
|
||||
dallService *dalle.Service,
|
||||
imageService *image.Service,
|
||||
userService *service.UserService) *FunctionHandler {
|
||||
return &FunctionHandler{
|
||||
BaseHandler: BaseHandler{
|
||||
@@ -48,7 +48,7 @@ func NewFunctionHandler(
|
||||
DB: db,
|
||||
},
|
||||
uploadManager: manager,
|
||||
dallService: dallService,
|
||||
imageService: imageService,
|
||||
userService: userService,
|
||||
}
|
||||
}
|
||||
@@ -61,7 +61,7 @@ func (h *FunctionHandler) RegisterRoutes() {
|
||||
// 需要用户授权的接口
|
||||
group.POST("weibo", h.WeiBo)
|
||||
group.POST("zaobao", h.ZaoBao)
|
||||
group.POST("dalle3", h.Dall3)
|
||||
group.POST("image3", h.Image3)
|
||||
}
|
||||
|
||||
type resVo struct {
|
||||
@@ -176,8 +176,8 @@ func (h *FunctionHandler) ZaoBao(c *gin.Context) {
|
||||
resp.SUCCESS(c, strings.Join(builder, "\n\n"))
|
||||
}
|
||||
|
||||
// Dall3 DallE3 AI 绘图
|
||||
func (h *FunctionHandler) Dall3(c *gin.Context) {
|
||||
// Image3 AI 图像生成
|
||||
func (h *FunctionHandler) Image3(c *gin.Context) {
|
||||
if err := h.checkAuth(c); err != nil {
|
||||
resp.ERROR(c, err.Error())
|
||||
return
|
||||
@@ -209,9 +209,9 @@ func (h *FunctionHandler) Dall3(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
// create dall task
|
||||
// create image task
|
||||
prompt := utils.InterfaceToString(params["prompt"])
|
||||
task := types.DallTask{
|
||||
task := types.ImageTask{
|
||||
UserId: user.Id,
|
||||
Prompt: prompt,
|
||||
ModelId: chatModel.Id,
|
||||
@@ -220,11 +220,11 @@ func (h *FunctionHandler) Dall3(c *gin.Context) {
|
||||
TranslateModelId: h.App.SysConfig.Base.AssistantModelId,
|
||||
Power: chatModel.Power,
|
||||
}
|
||||
job := model.DallJob{
|
||||
UserId: user.Id,
|
||||
Prompt: prompt,
|
||||
Power: chatModel.Power,
|
||||
TaskInfo: utils.JsonEncode(task),
|
||||
job := model.ImageJob{
|
||||
UserId: user.Id,
|
||||
Prompt: prompt,
|
||||
Power: chatModel.Power,
|
||||
Params: utils.JsonEncode(task),
|
||||
}
|
||||
err := h.DB.Create(&job).Error
|
||||
if err != nil {
|
||||
@@ -233,7 +233,7 @@ func (h *FunctionHandler) Dall3(c *gin.Context) {
|
||||
}
|
||||
|
||||
task.Id = job.Id
|
||||
content, err := h.dallService.Image(task, true)
|
||||
content, err := h.imageService.Image(task, true)
|
||||
if err != nil {
|
||||
resp.ERROR(c, "任务执行失败:"+err.Error())
|
||||
return
|
||||
|
||||
@@ -3,7 +3,7 @@ package handler
|
||||
// * +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
|
||||
// * Copyright 2023 The Geek-AI Authors. All rights reserved.
|
||||
// * Use of this source code is governed by a Apache-2.0 license
|
||||
// * that can be found in the LICENSE file.
|
||||
// * that can be found in LICENSE file.
|
||||
// * @Author yangjian102621@163.com
|
||||
// * +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
|
||||
|
||||
@@ -13,7 +13,7 @@ import (
|
||||
"geekai/core/middleware"
|
||||
"geekai/core/types"
|
||||
"geekai/service"
|
||||
"geekai/service/dalle"
|
||||
"geekai/service/image"
|
||||
"geekai/service/moderation"
|
||||
"geekai/service/oss"
|
||||
"geekai/store/model"
|
||||
@@ -25,17 +25,17 @@ import (
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type DallJobHandler struct {
|
||||
type ImageJobHandler struct {
|
||||
BaseHandler
|
||||
dallService *dalle.Service
|
||||
imageService *image.Service
|
||||
uploader *oss.UploaderManager
|
||||
userService *service.UserService
|
||||
moderationManager *moderation.ServiceManager
|
||||
}
|
||||
|
||||
func NewDallJobHandler(app *core.AppServer, db *gorm.DB, service *dalle.Service, manager *oss.UploaderManager, userService *service.UserService, moderationManager *moderation.ServiceManager) *DallJobHandler {
|
||||
return &DallJobHandler{
|
||||
dallService: service,
|
||||
func NewImageJobHandler(app *core.AppServer, db *gorm.DB, service *image.Service, manager *oss.UploaderManager, userService *service.UserService, moderationManager *moderation.ServiceManager) *ImageJobHandler {
|
||||
return &ImageJobHandler{
|
||||
imageService: service,
|
||||
uploader: manager,
|
||||
userService: userService,
|
||||
moderationManager: moderationManager,
|
||||
@@ -47,8 +47,8 @@ func NewDallJobHandler(app *core.AppServer, db *gorm.DB, service *dalle.Service,
|
||||
}
|
||||
|
||||
// RegisterRoutes 注册路由
|
||||
func (h *DallJobHandler) RegisterRoutes() {
|
||||
group := h.App.Engine.Group("/api/dall/")
|
||||
func (h *ImageJobHandler) RegisterRoutes() {
|
||||
group := h.App.Engine.Group("/api/image/")
|
||||
|
||||
// 公开接口,不需要授权
|
||||
group.GET("imgWall", h.ImgWall)
|
||||
@@ -65,8 +65,8 @@ func (h *DallJobHandler) RegisterRoutes() {
|
||||
}
|
||||
|
||||
// Image 创建一个绘画任务
|
||||
func (h *DallJobHandler) Image(c *gin.Context) {
|
||||
var data types.DallTask
|
||||
func (h *ImageJobHandler) Image(c *gin.Context) {
|
||||
var data types.ImageTask
|
||||
if err := c.ShouldBindJSON(&data); err != nil || data.Prompt == "" {
|
||||
resp.ERROR(c, types.InvalidArgs)
|
||||
return
|
||||
@@ -82,7 +82,7 @@ func (h *DallJobHandler) Image(c *gin.Context) {
|
||||
// 记录违规内容
|
||||
moderation := model.Moderation{
|
||||
UserId: h.GetLoginUserId(c),
|
||||
Source: types.ModerationSourceDalle,
|
||||
Source: types.ModerationSourceImage,
|
||||
Input: data.Prompt,
|
||||
Result: utils.JsonEncode(moderationResult),
|
||||
}
|
||||
@@ -114,7 +114,7 @@ func (h *DallJobHandler) Image(c *gin.Context) {
|
||||
|
||||
idValue, _ := c.Get(types.LoginUserID)
|
||||
userId := utils.IntValue(utils.InterfaceToString(idValue), 0)
|
||||
task := types.DallTask{
|
||||
task := types.ImageTask{
|
||||
UserId: uint(userId),
|
||||
ModelId: chatModel.Id,
|
||||
ModelName: chatModel.Name,
|
||||
@@ -126,11 +126,11 @@ func (h *DallJobHandler) Image(c *gin.Context) {
|
||||
TranslateModelId: h.App.SysConfig.Base.AssistantModelId,
|
||||
Power: chatModel.Power,
|
||||
}
|
||||
job := model.DallJob{
|
||||
UserId: uint(userId),
|
||||
Prompt: data.Prompt,
|
||||
Power: chatModel.Power,
|
||||
TaskInfo: utils.JsonEncode(task),
|
||||
job := model.ImageJob{
|
||||
UserId: uint(userId),
|
||||
Prompt: data.Prompt,
|
||||
Power: chatModel.Power,
|
||||
Params: utils.JsonEncode(task),
|
||||
}
|
||||
res := h.DB.Create(&job)
|
||||
if res.Error != nil {
|
||||
@@ -139,7 +139,7 @@ func (h *DallJobHandler) Image(c *gin.Context) {
|
||||
}
|
||||
|
||||
task.Id = job.Id
|
||||
h.dallService.PushTask(task)
|
||||
h.imageService.PushTask(task)
|
||||
|
||||
// 扣减算力
|
||||
err = h.userService.DecreasePower(user.Id, chatModel.Power, model.PowerLog{
|
||||
@@ -155,7 +155,7 @@ func (h *DallJobHandler) Image(c *gin.Context) {
|
||||
}
|
||||
|
||||
// ImgWall 照片墙
|
||||
func (h *DallJobHandler) ImgWall(c *gin.Context) {
|
||||
func (h *ImageJobHandler) ImgWall(c *gin.Context) {
|
||||
page := h.GetInt(c, "page", 0)
|
||||
pageSize := h.GetInt(c, "page_size", 0)
|
||||
err, jobs := h.getData(true, 0, page, pageSize, true)
|
||||
@@ -167,8 +167,8 @@ func (h *DallJobHandler) ImgWall(c *gin.Context) {
|
||||
resp.SUCCESS(c, jobs)
|
||||
}
|
||||
|
||||
// JobList 获取 SD 任务列表
|
||||
func (h *DallJobHandler) JobList(c *gin.Context) {
|
||||
// JobList 获取 Image 任务列表
|
||||
func (h *ImageJobHandler) JobList(c *gin.Context) {
|
||||
finish := h.GetBool(c, "finish")
|
||||
userId := h.GetLoginUserId(c)
|
||||
page := h.GetInt(c, "page", 0)
|
||||
@@ -185,7 +185,7 @@ func (h *DallJobHandler) JobList(c *gin.Context) {
|
||||
}
|
||||
|
||||
// JobList 获取任务列表
|
||||
func (h *DallJobHandler) getData(finish bool, userId uint, page int, pageSize int, publish bool) (error, vo.Page) {
|
||||
func (h *ImageJobHandler) getData(finish bool, userId uint, page int, pageSize int, publish bool) (error, vo.Page) {
|
||||
|
||||
session := h.DB.Session(&gorm.Session{})
|
||||
if finish {
|
||||
@@ -205,21 +205,22 @@ func (h *DallJobHandler) getData(finish bool, userId uint, page int, pageSize in
|
||||
}
|
||||
// 统计总数
|
||||
var total int64
|
||||
session.Model(&model.DallJob{}).Count(&total)
|
||||
session.Model(&model.ImageJob{}).Count(&total)
|
||||
|
||||
var items []model.DallJob
|
||||
var items []model.ImageJob
|
||||
res := session.Find(&items)
|
||||
if res.Error != nil {
|
||||
return res.Error, vo.Page{}
|
||||
}
|
||||
|
||||
var jobs = make([]vo.DallJob, 0)
|
||||
var jobs = make([]vo.ImageJob, 0)
|
||||
for _, item := range items {
|
||||
var job vo.DallJob
|
||||
var job vo.ImageJob
|
||||
err := utils.CopyObject(item, &job)
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
job.CreatedAt = item.CreatedAt.Unix()
|
||||
jobs = append(jobs, job)
|
||||
}
|
||||
|
||||
@@ -227,10 +228,10 @@ func (h *DallJobHandler) getData(finish bool, userId uint, page int, pageSize in
|
||||
}
|
||||
|
||||
// Remove remove task image
|
||||
func (h *DallJobHandler) Remove(c *gin.Context) {
|
||||
func (h *ImageJobHandler) Remove(c *gin.Context) {
|
||||
id := h.GetInt(c, "id", 0)
|
||||
userId := h.GetLoginUserId(c)
|
||||
var job model.DallJob
|
||||
var job model.ImageJob
|
||||
if res := h.DB.Where("id = ? AND user_id = ?", id, userId).First(&job); res.Error != nil {
|
||||
resp.ERROR(c, "记录不存在")
|
||||
return
|
||||
@@ -253,12 +254,12 @@ func (h *DallJobHandler) Remove(c *gin.Context) {
|
||||
}
|
||||
|
||||
// Publish 发布/取消发布图片到画廊显示
|
||||
func (h *DallJobHandler) Publish(c *gin.Context) {
|
||||
func (h *ImageJobHandler) Publish(c *gin.Context) {
|
||||
id := h.GetInt(c, "id", 0)
|
||||
userId := h.GetLoginUserId(c)
|
||||
action := h.GetBool(c, "action") // 发布动作,true => 发布,false => 取消分享
|
||||
|
||||
err := h.DB.Model(&model.DallJob{Id: uint(id), UserId: userId}).UpdateColumn("publish", action).Error
|
||||
err := h.DB.Model(&model.ImageJob{Id: uint(id), UserId: userId}).UpdateColumn("publish", action).Error
|
||||
if err != nil {
|
||||
resp.ERROR(c, err.Error())
|
||||
return
|
||||
@@ -267,7 +268,7 @@ func (h *DallJobHandler) Publish(c *gin.Context) {
|
||||
resp.SUCCESS(c)
|
||||
}
|
||||
|
||||
func (h *DallJobHandler) GetModels(c *gin.Context) {
|
||||
func (h *ImageJobHandler) GetModels(c *gin.Context) {
|
||||
var models []model.ChatModel
|
||||
err := h.DB.Where("type", "img").Where("enabled", true).Find(&models).Error
|
||||
if err != nil {
|
||||
+116
-6
@@ -63,6 +63,7 @@ func (h *MidJourneyHandler) RegisterRoutes() {
|
||||
group.POST("image", h.Image)
|
||||
group.POST("upscale", h.Upscale)
|
||||
group.POST("variation", h.Variation)
|
||||
group.POST("modal", h.Modal)
|
||||
group.GET("jobs", h.JobList)
|
||||
group.GET("remove", h.Remove)
|
||||
group.GET("publish", h.Publish)
|
||||
@@ -82,7 +83,31 @@ func (h *MidJourneyHandler) preCheck(c *gin.Context) bool {
|
||||
}
|
||||
|
||||
return true
|
||||
}
|
||||
|
||||
// preCheckPower 检查用户算力是否 >= required,不足时写 ERROR 并返回 false
|
||||
func (h *MidJourneyHandler) preCheckPower(c *gin.Context, required int) bool {
|
||||
user, err := h.GetLoginUser(c)
|
||||
if err != nil {
|
||||
resp.NotAuth(c)
|
||||
return false
|
||||
}
|
||||
if required <= 0 {
|
||||
required = h.App.SysConfig.Base.MjActionPower
|
||||
}
|
||||
if user.Power < required {
|
||||
resp.ERROR(c, "当前用户剩余算力不足以完成本次操作!")
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// mjActionPower 取分项算力,若未配置则回退 MjActionPower
|
||||
func mjActionPower(base int, fallback int) int {
|
||||
if base > 0 {
|
||||
return base
|
||||
}
|
||||
return fallback
|
||||
}
|
||||
|
||||
// Image 创建一个绘画任务
|
||||
@@ -109,7 +134,14 @@ func (h *MidJourneyHandler) Image(c *gin.Context) {
|
||||
resp.ERROR(c, types.InvalidArgs)
|
||||
return
|
||||
}
|
||||
if !h.preCheck(c) {
|
||||
// 按任务类型计算所需算力
|
||||
power := h.App.SysConfig.Base.MjPower
|
||||
if data.TaskType == types.TaskBlend.String() {
|
||||
power = mjActionPower(h.App.SysConfig.Base.MjBlendPower, h.App.SysConfig.Base.MjActionPower)
|
||||
} else if data.TaskType == types.TaskSwapFace.String() {
|
||||
power = mjActionPower(h.App.SysConfig.Base.MjSwapFacePower, h.App.SysConfig.Base.MjActionPower)
|
||||
}
|
||||
if !h.preCheckPower(c, power) {
|
||||
return
|
||||
}
|
||||
|
||||
@@ -215,7 +247,7 @@ func (h *MidJourneyHandler) Image(c *gin.Context) {
|
||||
TaskInfo: utils.JsonEncode(task),
|
||||
Progress: 0,
|
||||
Prompt: fmt.Sprintf("%s %s", data.Prompt, params),
|
||||
Power: h.App.SysConfig.Base.MjPower,
|
||||
Power: power,
|
||||
CreatedAt: time.Now(),
|
||||
}
|
||||
opt := "绘图"
|
||||
@@ -264,7 +296,8 @@ func (h *MidJourneyHandler) Upscale(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
if !h.preCheck(c) {
|
||||
power := mjActionPower(h.App.SysConfig.Base.MjUpscalePower, h.App.SysConfig.Base.MjActionPower)
|
||||
if !h.preCheckPower(c, power) {
|
||||
return
|
||||
}
|
||||
|
||||
@@ -286,7 +319,7 @@ func (h *MidJourneyHandler) Upscale(c *gin.Context) {
|
||||
TaskId: taskId,
|
||||
TaskInfo: utils.JsonEncode(task),
|
||||
Progress: 0,
|
||||
Power: h.App.SysConfig.Base.MjActionPower,
|
||||
Power: power,
|
||||
CreatedAt: time.Now(),
|
||||
}
|
||||
if res := h.DB.Create(&job); res.Error != nil || res.RowsAffected == 0 {
|
||||
@@ -319,7 +352,8 @@ func (h *MidJourneyHandler) Variation(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
if !h.preCheck(c) {
|
||||
power := mjActionPower(h.App.SysConfig.Base.MjUpscalePower, h.App.SysConfig.Base.MjActionPower)
|
||||
if !h.preCheckPower(c, power) {
|
||||
return
|
||||
}
|
||||
|
||||
@@ -342,7 +376,7 @@ func (h *MidJourneyHandler) Variation(c *gin.Context) {
|
||||
TaskId: taskId,
|
||||
TaskInfo: utils.JsonEncode(task),
|
||||
Progress: 0,
|
||||
Power: h.App.SysConfig.Base.MjActionPower,
|
||||
Power: power,
|
||||
CreatedAt: time.Now(),
|
||||
}
|
||||
if res := h.DB.Create(&job); res.Error != nil || res.RowsAffected == 0 {
|
||||
@@ -366,6 +400,81 @@ func (h *MidJourneyHandler) Variation(c *gin.Context) {
|
||||
resp.SUCCESS(c)
|
||||
}
|
||||
|
||||
// modalReq 局部重绘请求参数
|
||||
type modalReq struct {
|
||||
TaskId string `json:"task_id"` // 原图任务 ID(message_id)
|
||||
ChannelId string `json:"channel_id"` // 渠道 ID
|
||||
Prompt string `json:"prompt"` // 提示词
|
||||
MaskBase64 string `json:"mask_base64,omitempty"` // 蒙版 base64,可选
|
||||
}
|
||||
|
||||
// Modal 提交局部重绘(inpaint)
|
||||
func (h *MidJourneyHandler) Modal(c *gin.Context) {
|
||||
var data modalReq
|
||||
if err := c.ShouldBindJSON(&data); err != nil {
|
||||
resp.ERROR(c, types.InvalidArgs)
|
||||
return
|
||||
}
|
||||
if data.TaskId == "" || data.ChannelId == "" {
|
||||
resp.ERROR(c, "task_id 与 channel_id 必填")
|
||||
return
|
||||
}
|
||||
if data.Prompt == "" {
|
||||
resp.ERROR(c, "请填写局部重绘提示词")
|
||||
return
|
||||
}
|
||||
|
||||
power := mjActionPower(h.App.SysConfig.Base.MjModalPower, h.App.SysConfig.Base.MjActionPower)
|
||||
if !h.preCheckPower(c, power) {
|
||||
return
|
||||
}
|
||||
|
||||
idValue, _ := c.Get(types.LoginUserID)
|
||||
userId := utils.IntValue(utils.InterfaceToString(idValue), 0)
|
||||
taskId, _ := h.snowflake.Next(true)
|
||||
// 原图 message_id 必须传入 API,同时写入 TaskId/MessageId 避免序列化 omitempty 丢失
|
||||
task := types.MjTask{
|
||||
Type: types.TaskModal,
|
||||
UserId: userId,
|
||||
ChannelId: data.ChannelId,
|
||||
TaskId: data.TaskId,
|
||||
MessageId: data.TaskId,
|
||||
Prompt: data.Prompt,
|
||||
MaskBase64: data.MaskBase64,
|
||||
Mode: h.App.SysConfig.Base.MjMode,
|
||||
}
|
||||
job := model.MidJourneyJob{
|
||||
Type: types.TaskModal.String(),
|
||||
ChannelId: data.ChannelId,
|
||||
UserId: uint(userId),
|
||||
TaskId: taskId,
|
||||
TaskInfo: utils.JsonEncode(task),
|
||||
Progress: 0,
|
||||
Prompt: data.Prompt,
|
||||
Power: power,
|
||||
CreatedAt: time.Now(),
|
||||
}
|
||||
if res := h.DB.Create(&job); res.Error != nil || res.RowsAffected == 0 {
|
||||
resp.ERROR(c, "添加任务失败:"+res.Error.Error())
|
||||
return
|
||||
}
|
||||
|
||||
task.Id = job.Id
|
||||
h.mjService.PushTask(task)
|
||||
|
||||
err := h.userService.DecreasePower(job.UserId, job.Power, model.PowerLog{
|
||||
Type: types.PowerConsume,
|
||||
Model: "mid-journey",
|
||||
Remark: fmt.Sprintf("局部重绘操作,任务ID:%s", job.TaskId),
|
||||
})
|
||||
if err != nil {
|
||||
resp.ERROR(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
resp.SUCCESS(c)
|
||||
}
|
||||
|
||||
// ImgWall 照片墙
|
||||
func (h *MidJourneyHandler) ImgWall(c *gin.Context) {
|
||||
page := h.GetInt(c, "page", 0)
|
||||
@@ -432,6 +541,7 @@ func (h *MidJourneyHandler) getData(finish bool, userId uint, page int, pageSize
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
job.CreatedAt = item.CreatedAt.Unix()
|
||||
jobs = append(jobs, job)
|
||||
}
|
||||
return nil, vo.NewPage(total, page, pageSize, jobs)
|
||||
|
||||
@@ -216,7 +216,7 @@ func (h *PaymentHandler) CreateOrder(c *gin.Context) {
|
||||
data.Domain = h.config.WxPay.Domain
|
||||
}
|
||||
notifyURL = fmt.Sprintf("%s/api/payment/notify/wxpay", data.Domain)
|
||||
payURL, err = h.wxpayService.Pay(payment.PayRequest{
|
||||
params := payment.PayRequest{
|
||||
OutTradeNo: orderNo,
|
||||
TotalFee: fmt.Sprintf("%d", int(amount*100)),
|
||||
Subject: product.Name,
|
||||
@@ -224,7 +224,11 @@ func (h *PaymentHandler) CreateOrder(c *gin.Context) {
|
||||
ClientIP: c.ClientIP(),
|
||||
Device: data.Device,
|
||||
PayWay: payment.PayWayWX,
|
||||
})
|
||||
}
|
||||
if data.Device == "mobile" {
|
||||
params.OpenID = user.OpenId
|
||||
}
|
||||
payURL, err = h.wxpayService.Pay(params)
|
||||
if err != nil {
|
||||
resp.ERROR(c, err.Error())
|
||||
return
|
||||
|
||||
@@ -0,0 +1,547 @@
|
||||
package handler
|
||||
|
||||
// * +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
|
||||
// * Copyright 2023 The Geek-AI Authors. All rights reserved.
|
||||
// * Use of this source code is governed by a Apache-2.0 license
|
||||
// * that can be found in the LICENSE file.
|
||||
// * @Author yangjian102621@163.com
|
||||
// * +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"geekai/core"
|
||||
"geekai/core/middleware"
|
||||
"geekai/core/types"
|
||||
"geekai/service"
|
||||
"geekai/service/ppt"
|
||||
"geekai/utils"
|
||||
"geekai/utils/resp"
|
||||
"net/http"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// PPTTaskHandler 用户侧 PPT 生成任务处理器(薄层:参数、鉴权、调用 PptService)
|
||||
type PPTTaskHandler struct {
|
||||
BaseHandler
|
||||
snowflake *service.Snowflake
|
||||
pptService *ppt.PptService
|
||||
}
|
||||
|
||||
// NewPPTTaskHandler 创建 PPT 任务处理器
|
||||
func NewPPTTaskHandler(app *core.AppServer, db *gorm.DB, snowflake *service.Snowflake, pptService *ppt.PptService) *PPTTaskHandler {
|
||||
return &PPTTaskHandler{
|
||||
BaseHandler: BaseHandler{App: app, DB: db},
|
||||
snowflake: snowflake,
|
||||
pptService: pptService,
|
||||
}
|
||||
}
|
||||
|
||||
// RegisterRoutes 注册 PPT 任务相关路由
|
||||
func (h *PPTTaskHandler) RegisterRoutes() {
|
||||
group := h.App.Engine.Group("/api/v1/tasks")
|
||||
group.Use(middleware.UserAuthMiddleware(h.App.Config.Session.SecretKey, h.App.Redis))
|
||||
{
|
||||
group.POST("generate-slides", h.CreateTask)
|
||||
group.POST("generate-slides/from-file", h.CreateTaskFromFile)
|
||||
group.GET("", h.ListTasks) // 当前用户任务列表,须在 :task_id 前注册
|
||||
group.GET(":task_id/export", h.ExportTask)
|
||||
group.POST(":task_id/resume", h.ResumeTask)
|
||||
group.POST(":task_id/slides/:slide_index/edit-image", h.EditSlideImage)
|
||||
group.PATCH(":task_id/slides/:slide_index/active-image", h.SetActiveSlideImage)
|
||||
group.GET(":task_id", h.GetTask)
|
||||
group.DELETE(":task_id", h.DeleteTask)
|
||||
}
|
||||
}
|
||||
|
||||
type createPPTTaskRequest struct {
|
||||
Content string `json:"content"`
|
||||
Prompt string `json:"prompt"`
|
||||
Language string `json:"language"` // 如 zh-CN, en,约束分镜输出语言
|
||||
Pages int `json:"pages"` // 目标页数,0 表示用配置默认值
|
||||
Mode string `json:"mode"` // 生成模式:detailed=详细演示文稿,slides=演示用幻灯片,空则默认 slides
|
||||
}
|
||||
|
||||
type createPPTTaskFromFileRequest struct {
|
||||
FileURL string `json:"file_url"`
|
||||
Prompt string `json:"prompt"`
|
||||
Language string `json:"language"` // 如 zh-CN, en,约束分镜输出语言
|
||||
Pages int `json:"pages"` // 目标页数,0 表示用配置默认值
|
||||
Mode string `json:"mode"` // 生成模式:detailed=详细演示文稿,slides=演示用幻灯片,空则默认 slides
|
||||
}
|
||||
|
||||
// CreateTask 创建 PPT 生成任务(异步)
|
||||
func (h *PPTTaskHandler) CreateTask(c *gin.Context) {
|
||||
var req createPPTTaskRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil || req.Content == "" {
|
||||
resp.ERROR(c, types.InvalidArgs)
|
||||
return
|
||||
}
|
||||
|
||||
user, err := h.GetLoginUser(c)
|
||||
if err != nil {
|
||||
resp.NotAuth(c)
|
||||
return
|
||||
}
|
||||
|
||||
taskID, err := h.snowflake.Next(true)
|
||||
if err != nil {
|
||||
resp.ERROR(c, "生成任务ID失败:"+err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
task, cfg, err := h.pptService.BuildPendingTask(taskID, user.Id, int(user.Power), req.Content, req.Prompt, req.Language, req.Mode, req.Pages)
|
||||
if err != nil {
|
||||
if errors.Is(err, ppt.ErrInsufficientPower) {
|
||||
resp.ERROR(c, "当前用户算力不足以完成本次 PPT 生成任务!")
|
||||
return
|
||||
}
|
||||
resp.ERROR(c, "加载 PPT 配置或校验失败:"+err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
h.pptService.CreateTask(c.Request.Context(), task, cfg)
|
||||
go h.pptService.RunTask(context.Background(), task, cfg)
|
||||
|
||||
resp.SUCCESS(c, map[string]any{
|
||||
"task_id": taskID,
|
||||
"status": ppt.TaskStatusPending,
|
||||
})
|
||||
}
|
||||
|
||||
// CreateTaskFromFile 创建 PPT 生成任务(基于上传材料文本提炼)
|
||||
// 支持文件:PDF、Word(doc/docx)、TXT、Markdown(md/markdown)
|
||||
func (h *PPTTaskHandler) CreateTaskFromFile(c *gin.Context) {
|
||||
var req createPPTTaskFromFileRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil || strings.TrimSpace(req.FileURL) == "" {
|
||||
resp.ERROR(c, types.InvalidArgs)
|
||||
return
|
||||
}
|
||||
|
||||
fileURL := strings.TrimSpace(req.FileURL)
|
||||
|
||||
// 去掉 query,避免 ".pdf?xxx" 这种导致 ext 识别失败
|
||||
urlNoQuery := strings.Split(fileURL, "?")[0]
|
||||
ext := strings.ToLower(filepath.Ext(urlNoQuery))
|
||||
if ext == "" {
|
||||
resp.ERROR(c, "不支持的文件格式")
|
||||
return
|
||||
}
|
||||
|
||||
allowedExts := map[string]bool{
|
||||
".pdf": true,
|
||||
".doc": true,
|
||||
".docx": true,
|
||||
".txt": true,
|
||||
".md": true,
|
||||
".markdown": true,
|
||||
}
|
||||
if !allowedExts[ext] {
|
||||
resp.ERROR(c, "不支持的文件格式")
|
||||
return
|
||||
}
|
||||
|
||||
designPrompt := strings.TrimSpace(req.Prompt)
|
||||
language := strings.TrimSpace(req.Language)
|
||||
if language == "" {
|
||||
language = "zh-CN"
|
||||
}
|
||||
mode := strings.TrimSpace(req.Mode)
|
||||
|
||||
pages := req.Pages
|
||||
|
||||
user, err := h.GetLoginUser(c)
|
||||
if err != nil {
|
||||
resp.NotAuth(c)
|
||||
return
|
||||
}
|
||||
|
||||
taskID, err := h.snowflake.Next(true)
|
||||
if err != nil {
|
||||
resp.ERROR(c, "生成任务ID失败:"+err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
// 1) 下载文件并提取原始文本(Notebook 风格提炼前的输入)
|
||||
rawText, err := h.extractMaterialTextFromURL(c.Request.Context(), fileURL, ext)
|
||||
if err != nil || strings.TrimSpace(rawText) == "" {
|
||||
if err == nil {
|
||||
err = errors.New("empty extracted text")
|
||||
}
|
||||
resp.ERROR(c, "读取文件失败:"+err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
// 2) 校验算力与页数,组装待写入的 Task(先在提炼前做 power 检查)
|
||||
task, cfg, err := h.pptService.BuildPendingTask(taskID, user.Id, int(user.Power), rawText, designPrompt, language, mode, pages)
|
||||
if err != nil {
|
||||
if errors.Is(err, ppt.ErrInsufficientPower) {
|
||||
resp.ERROR(c, "当前用户算力不足以完成本次 PPT 生成任务!")
|
||||
return
|
||||
}
|
||||
resp.ERROR(c, "加载 PPT 配置或校验失败:"+err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
// 3) NotebookLM 风格提炼:rawText -> PPT 可用的 content(大纲/结构化要点)
|
||||
notebookContent, err := h.pptService.GenerateNotebookContent(c.Request.Context(), cfg, rawText, designPrompt, language)
|
||||
if err != nil || strings.TrimSpace(notebookContent) == "" {
|
||||
if err == nil {
|
||||
err = errors.New("empty notebook content")
|
||||
}
|
||||
resp.ERROR(c, "提炼材料失败:"+err.Error())
|
||||
return
|
||||
}
|
||||
task.Content = notebookContent
|
||||
|
||||
// 4) 落库 + 异步生成分镜与图片
|
||||
err = h.pptService.CreateTask(c.Request.Context(), task, cfg)
|
||||
if err != nil {
|
||||
resp.ERROR(c, "创建任务失败:"+err.Error())
|
||||
return
|
||||
}
|
||||
go h.pptService.RunTask(context.Background(), task, cfg)
|
||||
|
||||
resp.SUCCESS(c, map[string]any{
|
||||
"task_id": taskID,
|
||||
"status": ppt.TaskStatusPending,
|
||||
})
|
||||
}
|
||||
|
||||
func (h *PPTTaskHandler) extractMaterialTextFromURL(ctx context.Context, fileURL string, ext string) (string, error) {
|
||||
if strings.TrimSpace(fileURL) == "" {
|
||||
return "", errors.New("empty file url")
|
||||
}
|
||||
|
||||
switch ext {
|
||||
case ".txt", ".md", ".markdown":
|
||||
b, status, err := utils.FetchURLBytes(ctx, fileURL, "", 30*time.Second, 2, 8<<20)
|
||||
if err != nil {
|
||||
// status=0 时通常是请求阶段错误(例如 TLS 握手超时)
|
||||
return "", fmt.Errorf("download file failed: status=%d: %w", status, err)
|
||||
}
|
||||
return strings.TrimSpace(string(b)), nil
|
||||
default:
|
||||
// PDF/Word 等走 Tika
|
||||
return utils.ReadFileContent(fileURL, h.App.Config.TikaHost)
|
||||
}
|
||||
}
|
||||
|
||||
// ListTasks 当前用户的 PPT 任务列表(分页)
|
||||
func (h *PPTTaskHandler) ListTasks(c *gin.Context) {
|
||||
user, err := h.GetLoginUser(c)
|
||||
if err != nil {
|
||||
resp.NotAuth(c)
|
||||
return
|
||||
}
|
||||
|
||||
page := h.GetInt(c, "page", 1)
|
||||
pageSize := h.GetInt(c, "page_size", 20)
|
||||
if page <= 0 {
|
||||
page = 1
|
||||
}
|
||||
if pageSize <= 0 || pageSize > 100 {
|
||||
pageSize = 20
|
||||
}
|
||||
|
||||
slice, total := h.pptService.ListUserTasks(c.Request.Context(), user.Id, page, pageSize)
|
||||
|
||||
jobs := make([]map[string]any, 0, len(slice))
|
||||
for _, t := range slice {
|
||||
job := t.TaskSummaryMap()
|
||||
job["prompt"] = t.Prompt
|
||||
if t.ErrorMessage != "" {
|
||||
job["error_message"] = t.ErrorMessage
|
||||
}
|
||||
jobs = append(jobs, job)
|
||||
}
|
||||
|
||||
resp.SUCCESS(c, map[string]any{
|
||||
"jobs": jobs,
|
||||
"total": total,
|
||||
"page": page,
|
||||
"page_size": pageSize,
|
||||
})
|
||||
}
|
||||
|
||||
// ExportTask 导出任务幻灯片为 PDF 或 PPTX(按图片逐页)
|
||||
func (h *PPTTaskHandler) ExportTask(c *gin.Context) {
|
||||
taskID := strings.TrimSpace(c.Param("task_id"))
|
||||
if taskID == "" {
|
||||
resp.ERROR(c, types.InvalidArgs)
|
||||
return
|
||||
}
|
||||
ef, ok := ppt.ParseExportFormat(c.Query("format"))
|
||||
if !ok {
|
||||
resp.ERROR(c, "format 参数无效,支持 pdf 或 pptx")
|
||||
return
|
||||
}
|
||||
|
||||
user, err := h.GetLoginUser(c)
|
||||
if err != nil {
|
||||
resp.NotAuth(c)
|
||||
return
|
||||
}
|
||||
|
||||
task, exists := h.pptService.GetTask(taskID)
|
||||
if !exists {
|
||||
resp.ERROR(c, "任务不存在")
|
||||
return
|
||||
}
|
||||
if user.Id != task.UserID {
|
||||
resp.NotAuth(c)
|
||||
return
|
||||
}
|
||||
if task.Status != ppt.TaskStatusCompleted {
|
||||
resp.ERROR(c, "仅已完成任务可导出")
|
||||
return
|
||||
}
|
||||
|
||||
ossCfg := types.OSSConfig{}
|
||||
if h.App.SysConfig != nil {
|
||||
ossCfg = h.App.SysConfig.OSS
|
||||
}
|
||||
data, err := ppt.BuildExportBytes(c.Request.Context(), task.Slides, ef, ossCfg, h.App.Config)
|
||||
if err != nil {
|
||||
resp.ERROR(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
base := ppt.SanitizeExportBaseName(task.Title, task.TaskID)
|
||||
filename := base + ppt.ExportFileExt(ef)
|
||||
c.Header("Content-Disposition", ppt.ContentDispositionAttachment(filename))
|
||||
c.Data(http.StatusOK, ppt.ExportMimeType(ef), data)
|
||||
}
|
||||
|
||||
// ResumeTask 继续生成缺图页(POST /api/v1/tasks/:task_id/resume)
|
||||
func (h *PPTTaskHandler) ResumeTask(c *gin.Context) {
|
||||
taskID := strings.TrimSpace(c.Param("task_id"))
|
||||
if taskID == "" {
|
||||
resp.ERROR(c, types.InvalidArgs)
|
||||
return
|
||||
}
|
||||
|
||||
user, err := h.GetLoginUser(c)
|
||||
if err != nil {
|
||||
resp.NotAuth(c)
|
||||
return
|
||||
}
|
||||
|
||||
err = h.pptService.ResumeTask(c.Request.Context(), taskID, user.Id)
|
||||
if err != nil {
|
||||
if errors.Is(err, ppt.ErrPptTaskNotFound) {
|
||||
resp.ERROR(c, "任务不存在")
|
||||
return
|
||||
}
|
||||
if errors.Is(err, ppt.ErrPptTaskBusy) {
|
||||
resp.ERROR(c, "任务正在处理中,请稍后再试")
|
||||
return
|
||||
}
|
||||
if errors.Is(err, ppt.ErrPptTaskNotResumable) {
|
||||
resp.ERROR(c, "当前任务无法继续生成(已完成或分镜数据不完整)")
|
||||
return
|
||||
}
|
||||
if errors.Is(err, ppt.ErrInsufficientPower) {
|
||||
resp.ERROR(c, "当前用户算力不足以完成剩余图片生成")
|
||||
return
|
||||
}
|
||||
resp.ERROR(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
resp.SUCCESS(c, map[string]any{
|
||||
"task_id": taskID,
|
||||
"status": ppt.TaskStatusProcessing,
|
||||
})
|
||||
}
|
||||
|
||||
type editSlideImageRequest struct {
|
||||
Prompt string `json:"prompt"`
|
||||
}
|
||||
|
||||
type activeSlideImageRequest struct {
|
||||
VersionIndex int `json:"version_index"`
|
||||
}
|
||||
|
||||
// EditSlideImage 图生图编辑当前页(基于激活图)
|
||||
func (h *PPTTaskHandler) EditSlideImage(c *gin.Context) {
|
||||
taskID := strings.TrimSpace(c.Param("task_id"))
|
||||
slideIndexStr := strings.TrimSpace(c.Param("slide_index"))
|
||||
if taskID == "" || slideIndexStr == "" {
|
||||
resp.ERROR(c, types.InvalidArgs)
|
||||
return
|
||||
}
|
||||
slideIndex, err := strconv.Atoi(slideIndexStr)
|
||||
if err != nil || slideIndex < 1 {
|
||||
resp.ERROR(c, types.InvalidArgs)
|
||||
return
|
||||
}
|
||||
|
||||
var req editSlideImageRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
resp.ERROR(c, types.InvalidArgs)
|
||||
return
|
||||
}
|
||||
|
||||
user, err := h.GetLoginUser(c)
|
||||
if err != nil {
|
||||
resp.NotAuth(c)
|
||||
return
|
||||
}
|
||||
|
||||
ossCfg := types.OSSConfig{}
|
||||
if h.App.SysConfig != nil {
|
||||
ossCfg = h.App.SysConfig.OSS
|
||||
}
|
||||
|
||||
slides, err := h.pptService.EditSlideImage(c.Request.Context(), taskID, user.Id, slideIndex, req.Prompt, ossCfg, h.App.Config)
|
||||
if err != nil {
|
||||
if errors.Is(err, ppt.ErrPptTaskNotFound) {
|
||||
resp.ERROR(c, "任务不存在")
|
||||
return
|
||||
}
|
||||
if errors.Is(err, ppt.ErrPptSlideNotFound) {
|
||||
resp.ERROR(c, "幻灯片不存在")
|
||||
return
|
||||
}
|
||||
if errors.Is(err, ppt.ErrPptSlideNoImage) {
|
||||
resp.ERROR(c, "该页暂无配图,无法编辑")
|
||||
return
|
||||
}
|
||||
if errors.Is(err, ppt.ErrInsufficientPower) {
|
||||
resp.ERROR(c, "当前用户算力不足以完成本次编辑")
|
||||
return
|
||||
}
|
||||
resp.ERROR(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
resp.SUCCESS(c, map[string]any{"slides": slides})
|
||||
}
|
||||
|
||||
// SetActiveSlideImage 切换当前页激活的历史版本
|
||||
func (h *PPTTaskHandler) SetActiveSlideImage(c *gin.Context) {
|
||||
taskID := strings.TrimSpace(c.Param("task_id"))
|
||||
slideIndexStr := strings.TrimSpace(c.Param("slide_index"))
|
||||
if taskID == "" || slideIndexStr == "" {
|
||||
resp.ERROR(c, types.InvalidArgs)
|
||||
return
|
||||
}
|
||||
slideIndex, err := strconv.Atoi(slideIndexStr)
|
||||
if err != nil || slideIndex < 1 {
|
||||
resp.ERROR(c, types.InvalidArgs)
|
||||
return
|
||||
}
|
||||
|
||||
var req activeSlideImageRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
resp.ERROR(c, types.InvalidArgs)
|
||||
return
|
||||
}
|
||||
|
||||
user, err := h.GetLoginUser(c)
|
||||
if err != nil {
|
||||
resp.NotAuth(c)
|
||||
return
|
||||
}
|
||||
|
||||
slides, err := h.pptService.SetActiveSlideVersion(taskID, user.Id, slideIndex, req.VersionIndex)
|
||||
if err != nil {
|
||||
if errors.Is(err, ppt.ErrPptTaskNotFound) {
|
||||
resp.ERROR(c, "任务不存在")
|
||||
return
|
||||
}
|
||||
if errors.Is(err, ppt.ErrPptSlideNotFound) {
|
||||
resp.ERROR(c, "幻灯片不存在")
|
||||
return
|
||||
}
|
||||
if errors.Is(err, ppt.ErrPptInvalidVersionIndex) {
|
||||
resp.ERROR(c, "无效的历史版本序号")
|
||||
return
|
||||
}
|
||||
resp.ERROR(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
resp.SUCCESS(c, map[string]any{"slides": slides})
|
||||
}
|
||||
|
||||
// GetTask 查询 PPT 任务进度
|
||||
func (h *PPTTaskHandler) GetTask(c *gin.Context) {
|
||||
taskId := c.Param("task_id")
|
||||
if taskId == "" {
|
||||
resp.ERROR(c, types.InvalidArgs)
|
||||
return
|
||||
}
|
||||
|
||||
task, ok := h.pptService.GetTask(taskId)
|
||||
if !ok {
|
||||
resp.ERROR(c, "任务不存在")
|
||||
return
|
||||
}
|
||||
|
||||
user, err := h.GetLoginUser(c)
|
||||
if err != nil || user.Id != task.UserID {
|
||||
resp.NotAuth(c)
|
||||
return
|
||||
}
|
||||
|
||||
percentage := 0
|
||||
if task.Total > 0 {
|
||||
percentage = int(float64(task.Completed) / float64(task.Total) * 100)
|
||||
}
|
||||
|
||||
resp.SUCCESS(c, map[string]any{
|
||||
"task_id": task.TaskID,
|
||||
"status": task.Status,
|
||||
"progress": map[string]any{
|
||||
"total_slides": task.Total,
|
||||
"completed_slides": task.Completed,
|
||||
"percentage": percentage,
|
||||
},
|
||||
"slides": task.Slides,
|
||||
"error_message": task.ErrorMessage,
|
||||
"content": task.Content,
|
||||
"prompt": task.Prompt,
|
||||
"title": task.Title,
|
||||
"thumb": task.Thumb,
|
||||
})
|
||||
}
|
||||
|
||||
// DeleteTask 删除用户 PPT 任务(仅允许 completed / failed),同时删除该任务生成的图片对象。
|
||||
func (h *PPTTaskHandler) DeleteTask(c *gin.Context) {
|
||||
taskID := c.Param("task_id")
|
||||
if taskID == "" {
|
||||
resp.ERROR(c, types.InvalidArgs)
|
||||
return
|
||||
}
|
||||
|
||||
user, err := h.GetLoginUser(c)
|
||||
if err != nil {
|
||||
resp.NotAuth(c)
|
||||
return
|
||||
}
|
||||
|
||||
if err := h.pptService.DeleteTask(taskID, user.Id); err != nil {
|
||||
if errors.Is(err, ppt.ErrPptTaskNotFound) {
|
||||
resp.ERROR(c, "任务不存在")
|
||||
return
|
||||
}
|
||||
if errors.Is(err, ppt.ErrPptTaskNotDeletable) {
|
||||
resp.ERROR(c, "仅允许删除已完成/已失败任务")
|
||||
return
|
||||
}
|
||||
// 前端会统一在文案里拼接“删除失败:”,这里避免重复。
|
||||
resp.ERROR(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
resp.SUCCESS(c, gin.H{"message": "删除成功"})
|
||||
}
|
||||
@@ -179,7 +179,7 @@ func (h *RealtimeHandler) VoiceChat(c *gin.Context) {
|
||||
}
|
||||
apiURL := fmt.Sprintf("%s/v1/chat/completions", apiKey.ApiURL)
|
||||
logger.Infof("Sending %s request, API KEY:%s, PROXY: %s, Model: %s", apiKey.ApiURL, apiURL, apiKey.ProxyURL, "advanced-voice")
|
||||
r, err := client.R().SetHeader("Body-Type", "application/json").
|
||||
r, err := client.R().SetHeader("Content-Type", "application/json").
|
||||
SetHeader("Authorization", "Bearer "+apiKey.Value).
|
||||
SetBody(types.ApiRequest{
|
||||
Model: "advanced-voice",
|
||||
@@ -221,11 +221,12 @@ func (h *RealtimeHandler) VoiceChat(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
logger.Infof("Response: %v", response.Choices[0].Message.Content)
|
||||
replyText := utils.NormalizeAssistantContent(response.Choices[0].Message.Content)
|
||||
logger.Infof("Response: %v", replyText)
|
||||
|
||||
// 提取链接
|
||||
re := regexp.MustCompile(`\[(.*?)\]\((.*?)\)`)
|
||||
links := re.FindAllStringSubmatch(response.Choices[0].Message.Content, -1)
|
||||
links := re.FindAllStringSubmatch(replyText, -1)
|
||||
var url = ""
|
||||
if len(links) > 0 {
|
||||
url = links[0][2]
|
||||
|
||||
@@ -1,328 +0,0 @@
|
||||
package handler
|
||||
|
||||
// * +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
|
||||
// * Copyright 2023 The Geek-AI Authors. All rights reserved.
|
||||
// * Use of this source code is governed by a Apache-2.0 license
|
||||
// * that can be found in the LICENSE file.
|
||||
// * @Author yangjian102621@163.com
|
||||
// * +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"geekai/core"
|
||||
"geekai/core/middleware"
|
||||
"geekai/core/types"
|
||||
"geekai/service"
|
||||
"geekai/service/moderation"
|
||||
"geekai/service/oss"
|
||||
"geekai/service/sd"
|
||||
"geekai/store"
|
||||
"geekai/store/model"
|
||||
"geekai/store/vo"
|
||||
"geekai/utils"
|
||||
"geekai/utils/resp"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/go-redis/redis/v8"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type SdJobHandler struct {
|
||||
BaseHandler
|
||||
redis *redis.Client
|
||||
sdService *sd.Service
|
||||
uploader *oss.UploaderManager
|
||||
snowflake *service.Snowflake
|
||||
leveldb *store.LevelDB
|
||||
userService *service.UserService
|
||||
moderationManager *moderation.ServiceManager
|
||||
}
|
||||
|
||||
func NewSdJobHandler(app *core.AppServer,
|
||||
db *gorm.DB,
|
||||
service *sd.Service,
|
||||
manager *oss.UploaderManager,
|
||||
snowflake *service.Snowflake,
|
||||
userService *service.UserService,
|
||||
levelDB *store.LevelDB,
|
||||
moderationManager *moderation.ServiceManager) *SdJobHandler {
|
||||
return &SdJobHandler{
|
||||
sdService: service,
|
||||
uploader: manager,
|
||||
snowflake: snowflake,
|
||||
leveldb: levelDB,
|
||||
userService: userService,
|
||||
moderationManager: moderationManager,
|
||||
BaseHandler: BaseHandler{
|
||||
App: app,
|
||||
DB: db,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// RegisterRoutes 注册路由
|
||||
func (h *SdJobHandler) RegisterRoutes() {
|
||||
group := h.App.Engine.Group("/api/sd/")
|
||||
|
||||
// 公开接口,不需要授权
|
||||
group.GET("imgWall", h.ImgWall)
|
||||
|
||||
// 需要用户授权的接口
|
||||
group.Use(middleware.UserAuthMiddleware(h.App.Config.Session.SecretKey, h.App.Redis))
|
||||
{
|
||||
group.POST("image", h.Image)
|
||||
group.GET("jobs", h.JobList)
|
||||
group.GET("remove", h.Remove)
|
||||
group.GET("publish", h.Publish)
|
||||
}
|
||||
}
|
||||
|
||||
func (h *SdJobHandler) preCheck(c *gin.Context) bool {
|
||||
user, err := h.GetLoginUser(c)
|
||||
if err != nil {
|
||||
resp.NotAuth(c)
|
||||
return false
|
||||
}
|
||||
|
||||
if user.Power < h.App.SysConfig.Base.SdPower {
|
||||
resp.ERROR(c, "当前用户剩余算力不足以完成本次绘画!")
|
||||
return false
|
||||
}
|
||||
|
||||
return true
|
||||
|
||||
}
|
||||
|
||||
// Image 创建一个绘画任务
|
||||
func (h *SdJobHandler) Image(c *gin.Context) {
|
||||
if !h.preCheck(c) {
|
||||
return
|
||||
}
|
||||
|
||||
var data types.SdTaskParams
|
||||
if err := c.ShouldBindJSON(&data); err != nil || data.Prompt == "" {
|
||||
resp.ERROR(c, types.InvalidArgs)
|
||||
return
|
||||
}
|
||||
|
||||
if h.App.SysConfig.Moderation.Enable {
|
||||
moderationResult, err := h.moderationManager.GetService().Moderate(data.Prompt)
|
||||
if err != nil {
|
||||
logger.Error("failed to moderate content: ", err)
|
||||
}
|
||||
if moderationResult.Flagged {
|
||||
// 记录违规内容
|
||||
moderation := model.Moderation{
|
||||
UserId: h.GetLoginUserId(c),
|
||||
Source: types.ModerationSourceSD,
|
||||
Input: data.Prompt,
|
||||
Result: utils.JsonEncode(moderationResult),
|
||||
}
|
||||
err = h.DB.Create(&moderation).Error
|
||||
if err != nil {
|
||||
logger.Error("failed to save moderation: ", err)
|
||||
}
|
||||
resp.ERROR(c, "当前创作内容包含敏感词,请重新输入!")
|
||||
return
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
if data.Width <= 0 {
|
||||
data.Width = 512
|
||||
}
|
||||
if data.Height <= 0 {
|
||||
data.Height = 512
|
||||
}
|
||||
if data.CfgScale <= 0 {
|
||||
data.CfgScale = 7
|
||||
}
|
||||
if data.Seed == 0 {
|
||||
data.Seed = -1
|
||||
}
|
||||
if data.Steps <= 0 {
|
||||
data.Steps = 20
|
||||
}
|
||||
if data.Sampler == "" {
|
||||
data.Sampler = "Euler a"
|
||||
}
|
||||
|
||||
idValue, _ := c.Get(types.LoginUserID)
|
||||
userId := utils.IntValue(utils.InterfaceToString(idValue), 0)
|
||||
taskId, err := h.snowflake.Next(true)
|
||||
if err != nil {
|
||||
resp.ERROR(c, "error with generate task id: "+err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
task := types.SdTask{
|
||||
Type: types.TaskImage,
|
||||
Params: types.SdTaskParams{
|
||||
TaskId: taskId,
|
||||
Prompt: data.Prompt,
|
||||
NegPrompt: data.NegPrompt,
|
||||
Steps: data.Steps,
|
||||
Sampler: data.Sampler,
|
||||
FaceFix: data.FaceFix,
|
||||
CfgScale: data.CfgScale,
|
||||
Seed: data.Seed,
|
||||
Height: data.Height,
|
||||
Width: data.Width,
|
||||
HdFix: data.HdFix,
|
||||
HdRedrawRate: data.HdRedrawRate,
|
||||
HdScale: data.HdScale,
|
||||
HdScaleAlg: data.HdScaleAlg,
|
||||
HdSteps: data.HdSteps,
|
||||
},
|
||||
UserId: userId,
|
||||
TranslateModelId: h.App.SysConfig.Base.AssistantModelId,
|
||||
}
|
||||
|
||||
job := model.SdJob{
|
||||
UserId: uint(userId),
|
||||
Type: types.TaskImage.String(),
|
||||
TaskId: taskId,
|
||||
Params: utils.JsonEncode(task.Params),
|
||||
TaskInfo: utils.JsonEncode(task),
|
||||
Prompt: data.Prompt,
|
||||
Progress: 0,
|
||||
Power: h.App.SysConfig.Base.SdPower,
|
||||
CreatedAt: time.Now(),
|
||||
}
|
||||
res := h.DB.Create(&job)
|
||||
if res.Error != nil {
|
||||
resp.ERROR(c, "error with save job: "+res.Error.Error())
|
||||
return
|
||||
}
|
||||
|
||||
task.Id = int(job.Id)
|
||||
h.sdService.PushTask(task)
|
||||
|
||||
// update user's power
|
||||
err = h.userService.DecreasePower(job.UserId, job.Power, model.PowerLog{
|
||||
Type: types.PowerConsume,
|
||||
Model: "stable-diffusion",
|
||||
Remark: fmt.Sprintf("绘图操作,任务ID:%s", job.TaskId),
|
||||
})
|
||||
if err != nil {
|
||||
resp.ERROR(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
resp.SUCCESS(c)
|
||||
}
|
||||
|
||||
// ImgWall 照片墙
|
||||
func (h *SdJobHandler) ImgWall(c *gin.Context) {
|
||||
page := h.GetInt(c, "page", 0)
|
||||
pageSize := h.GetInt(c, "page_size", 0)
|
||||
err, jobs := h.getData(true, 0, page, pageSize, true)
|
||||
if err != nil {
|
||||
resp.ERROR(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
resp.SUCCESS(c, jobs)
|
||||
}
|
||||
|
||||
// JobList 获取 SD 任务列表
|
||||
func (h *SdJobHandler) JobList(c *gin.Context) {
|
||||
finish := h.GetBool(c, "finish")
|
||||
userId := h.GetLoginUserId(c)
|
||||
page := h.GetInt(c, "page", 0)
|
||||
pageSize := h.GetInt(c, "page_size", 0)
|
||||
publish := h.GetBool(c, "publish")
|
||||
|
||||
err, jobs := h.getData(finish, userId, page, pageSize, publish)
|
||||
if err != nil {
|
||||
resp.ERROR(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
resp.SUCCESS(c, jobs)
|
||||
}
|
||||
|
||||
// JobList 获取 MJ 任务列表
|
||||
func (h *SdJobHandler) getData(finish bool, userId uint, page int, pageSize int, publish bool) (error, vo.Page) {
|
||||
|
||||
session := h.DB.Session(&gorm.Session{})
|
||||
if finish {
|
||||
session = session.Where("progress >= ?", 100).Order("id DESC")
|
||||
} else {
|
||||
session = session.Where("progress < ?", 100).Order("id ASC")
|
||||
}
|
||||
if userId > 0 {
|
||||
session = session.Where("user_id = ?", userId)
|
||||
}
|
||||
if publish {
|
||||
session = session.Where("publish", publish)
|
||||
}
|
||||
if page > 0 && pageSize > 0 {
|
||||
offset := (page - 1) * pageSize
|
||||
session = session.Offset(offset).Limit(pageSize)
|
||||
}
|
||||
|
||||
// 统计总数
|
||||
var total int64
|
||||
session.Model(&model.SdJob{}).Count(&total)
|
||||
|
||||
var items []model.SdJob
|
||||
res := session.Find(&items)
|
||||
if res.Error != nil {
|
||||
return res.Error, vo.Page{}
|
||||
}
|
||||
|
||||
var jobs = make([]vo.SdJob, 0)
|
||||
for _, item := range items {
|
||||
var job vo.SdJob
|
||||
err := utils.CopyObject(item, &job)
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
jobs = append(jobs, job)
|
||||
}
|
||||
|
||||
return nil, vo.NewPage(total, page, pageSize, jobs)
|
||||
}
|
||||
|
||||
// Remove remove task image
|
||||
func (h *SdJobHandler) Remove(c *gin.Context) {
|
||||
id := h.GetInt(c, "id", 0)
|
||||
userId := h.GetLoginUserId(c)
|
||||
var job model.SdJob
|
||||
if res := h.DB.Where("id = ? AND user_id = ?", id, userId).First(&job); res.Error != nil {
|
||||
resp.ERROR(c, "记录不存在")
|
||||
return
|
||||
}
|
||||
|
||||
// 删除任务
|
||||
err := h.DB.Delete(&job).Error
|
||||
if err != nil {
|
||||
resp.ERROR(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
// remove image
|
||||
err = h.uploader.GetUploadHandler().Delete(job.ImgURL)
|
||||
if err != nil {
|
||||
logger.Error("remove image failed: ", err)
|
||||
}
|
||||
|
||||
resp.SUCCESS(c)
|
||||
}
|
||||
|
||||
// Publish 发布/取消发布图片到画廊显示
|
||||
func (h *SdJobHandler) Publish(c *gin.Context) {
|
||||
id := h.GetInt(c, "id", 0)
|
||||
userId := h.GetLoginUserId(c)
|
||||
action := h.GetBool(c, "action") // 发布动作,true => 发布,false => 取消分享
|
||||
|
||||
err := h.DB.Model(&model.SdJob{Id: uint(id), UserId: uint(userId)}).UpdateColumn("publish", action).Error
|
||||
if err != nil {
|
||||
resp.ERROR(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
resp.SUCCESS(c)
|
||||
}
|
||||
+22
-18
@@ -125,9 +125,9 @@ func (h *SunoHandler) Create(c *gin.Context) {
|
||||
if data.SongId != "" && data.Type == 3 {
|
||||
var song model.SunoJob
|
||||
if err := h.DB.Where("song_id = ?", data.SongId).First(&song).Error; err == nil {
|
||||
data.Instrumental = song.Instrumental
|
||||
data.Model = song.ModelName
|
||||
data.Tags = song.Tags
|
||||
data.Instrumental = song.Params.Instrumental
|
||||
data.Model = song.Params.Model
|
||||
data.Tags = song.Params.Tags
|
||||
}
|
||||
// 拼接歌词
|
||||
var refSong model.SunoJob
|
||||
@@ -153,22 +153,26 @@ func (h *SunoHandler) Create(c *gin.Context) {
|
||||
|
||||
// 插入数据库
|
||||
job := model.SunoJob{
|
||||
UserId: uint(task.UserId),
|
||||
Prompt: data.Prompt,
|
||||
Instrumental: data.Instrumental,
|
||||
ModelName: data.Model,
|
||||
TaskInfo: utils.JsonEncode(task),
|
||||
Tags: data.Tags,
|
||||
Title: data.Title,
|
||||
Type: data.Type,
|
||||
RefSongId: data.RefSongId,
|
||||
RefTaskId: data.RefTaskId,
|
||||
ExtendSecs: data.ExtendSecs,
|
||||
Power: h.App.SysConfig.Base.SunoPower,
|
||||
SongId: utils.RandString(32),
|
||||
UserId: uint(task.UserId),
|
||||
Prompt: data.Prompt,
|
||||
Params: vo.SunoParam{
|
||||
Prompt: data.Prompt,
|
||||
Instrumental: data.Instrumental,
|
||||
Tags: data.Tags,
|
||||
ExtendSecs: data.ExtendSecs,
|
||||
Lyrics: data.Lyrics,
|
||||
Model: data.Model,
|
||||
},
|
||||
Title: data.Title,
|
||||
Type: data.Type,
|
||||
RefSongId: data.RefSongId,
|
||||
RefTaskId: data.RefTaskId,
|
||||
Power: h.App.SysConfig.Base.SunoPower,
|
||||
SongId: utils.RandString(32),
|
||||
}
|
||||
if data.Lyrics != "" {
|
||||
job.Prompt = data.Lyrics
|
||||
job.Params.Prompt = data.Lyrics
|
||||
}
|
||||
tx := h.DB.Create(&job)
|
||||
if tx.Error != nil {
|
||||
@@ -183,8 +187,8 @@ func (h *SunoHandler) Create(c *gin.Context) {
|
||||
// update user's power
|
||||
err = h.userService.DecreasePower(job.UserId, job.Power, model.PowerLog{
|
||||
Type: types.PowerConsume,
|
||||
Model: job.ModelName,
|
||||
Remark: fmt.Sprintf("Suno 文生歌曲,%s", job.ModelName),
|
||||
Model: job.Params.Model,
|
||||
Remark: fmt.Sprintf("Suno 文生歌曲,%s", job.Params.Model),
|
||||
CreatedAt: time.Now(),
|
||||
})
|
||||
if err != nil {
|
||||
|
||||
+139
-9
@@ -39,6 +39,7 @@ type UserHandler struct {
|
||||
userService *service.UserService
|
||||
wxLoginService *service.WxLoginService
|
||||
ipSearcher *xdb.Searcher
|
||||
wechatService *service.WxGzhService
|
||||
}
|
||||
|
||||
func NewUserHandler(
|
||||
@@ -50,6 +51,7 @@ func NewUserHandler(
|
||||
captcha *service.CaptchaService,
|
||||
userService *service.UserService,
|
||||
wxLoginService *service.WxLoginService,
|
||||
wechatService *service.WxGzhService,
|
||||
ipSearcher *xdb.Searcher) *UserHandler {
|
||||
return &UserHandler{
|
||||
BaseHandler: BaseHandler{DB: db, App: app},
|
||||
@@ -59,6 +61,7 @@ func NewUserHandler(
|
||||
captchaService: captcha,
|
||||
userService: userService,
|
||||
wxLoginService: wxLoginService,
|
||||
wechatService: wechatService,
|
||||
ipSearcher: ipSearcher,
|
||||
}
|
||||
}
|
||||
@@ -75,6 +78,7 @@ func (h *UserHandler) RegisterRoutes() {
|
||||
group.POST("login/callback", h.WxLoginCallback)
|
||||
group.GET("login/status", h.GetWxLoginState)
|
||||
group.GET("logout", h.Logout)
|
||||
group.POST("wxAuthLogin", h.WxAuthLogin)
|
||||
|
||||
// 需要用户授权的接口
|
||||
group.Use(middleware.UserAuthMiddleware(h.App.Config.Session.SecretKey, h.App.Redis))
|
||||
@@ -319,11 +323,22 @@ func (h *UserHandler) GetWxLoginState(c *gin.Context) {
|
||||
func (h *UserHandler) createNewUser(user model.User, code string) (model.User, error) {
|
||||
if user.OpenId != "" {
|
||||
user.Platform = "wechat"
|
||||
user.Nickname = fmt.Sprintf("微信用户@%d", utils.RandomNumber(6))
|
||||
user.Username = fmt.Sprintf("wx@%d", utils.RandomNumber(8))
|
||||
user.Password = "geekai123"
|
||||
// 如果未设置昵称,则生成默认昵称
|
||||
if user.Nickname == "" {
|
||||
user.Nickname = fmt.Sprintf("微信用户@%d", utils.RandomNumber(6))
|
||||
}
|
||||
// 如果未设置用户名,则生成默认用户名
|
||||
if user.Username == "" {
|
||||
user.Username = fmt.Sprintf("wx@%d", utils.RandomNumber(8))
|
||||
}
|
||||
// 如果未设置密码,则生成默认密码
|
||||
if user.Password == "" {
|
||||
user.Password = "geekai123"
|
||||
}
|
||||
} else {
|
||||
user.Nickname = fmt.Sprintf("用户@%d", utils.RandomNumber(6))
|
||||
if user.Nickname == "" {
|
||||
user.Nickname = fmt.Sprintf("用户@%d", utils.RandomNumber(6))
|
||||
}
|
||||
if user.Username == "" || user.Password == "" {
|
||||
return user, fmt.Errorf("用户名或密码不能为空")
|
||||
}
|
||||
@@ -332,9 +347,11 @@ func (h *UserHandler) createNewUser(user model.User, code string) (model.User, e
|
||||
salt := utils.RandString(8)
|
||||
user.Salt = salt
|
||||
user.Password = utils.GenPassword(user.Password, salt)
|
||||
user.Avatar = "/images/avatar/user.png"
|
||||
// 如果未设置头像,则使用默认头像
|
||||
if user.Avatar == "" {
|
||||
user.Avatar = "/images/avatar/user.png"
|
||||
}
|
||||
user.Status = true
|
||||
user.ChatRoles = utils.JsonEncode([]string{"gpt"})
|
||||
user.ChatConfig = "{}"
|
||||
user.ChatModels = "{}"
|
||||
user.Power = h.App.SysConfig.Base.InitPower
|
||||
@@ -471,6 +488,20 @@ func (h *UserHandler) Session(c *gin.Context) {
|
||||
h.DB.Model(&user).UpdateColumn("vip", false)
|
||||
}
|
||||
userVo.Id = user.Id
|
||||
// 工作区应用 ID 列表(历史可能为 key 数组,仅解析数字 ID)
|
||||
if user.ChatRoles != "" {
|
||||
var raw []interface{}
|
||||
if utils.JsonDecode(user.ChatRoles, &raw) == nil {
|
||||
for _, v := range raw {
|
||||
if n, ok := v.(float64); ok && n >= 0 {
|
||||
userVo.ChatRoles = append(userVo.ChatRoles, uint(n))
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
if userVo.ChatRoles == nil {
|
||||
userVo.ChatRoles = []uint{}
|
||||
}
|
||||
resp.SUCCESS(c, userVo)
|
||||
|
||||
}
|
||||
@@ -483,6 +514,7 @@ type userProfile struct {
|
||||
Power int `json:"power"`
|
||||
ExpiredTime int64 `json:"expired_time"`
|
||||
Vip bool `json:"vip"`
|
||||
GemIds []uint `json:"gem_ids"`
|
||||
}
|
||||
|
||||
func (h *UserHandler) Profile(c *gin.Context) {
|
||||
@@ -502,6 +534,19 @@ func (h *UserHandler) Profile(c *gin.Context) {
|
||||
}
|
||||
|
||||
profile.Id = user.Id
|
||||
if user.GemIds != "" {
|
||||
var raw []interface{}
|
||||
if utils.JsonDecode(user.GemIds, &raw) == nil {
|
||||
for _, v := range raw {
|
||||
if n, ok := v.(float64); ok {
|
||||
profile.GemIds = append(profile.GemIds, uint(n))
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
if profile.GemIds == nil {
|
||||
profile.GemIds = []uint{}
|
||||
}
|
||||
resp.SUCCESS(c, profile)
|
||||
}
|
||||
|
||||
@@ -520,6 +565,12 @@ func (h *UserHandler) ProfileUpdate(c *gin.Context) {
|
||||
h.DB.First(&user, user.Id)
|
||||
user.Avatar = data.Avatar
|
||||
user.Nickname = data.Nickname
|
||||
if data.GemIds != nil {
|
||||
if len(data.GemIds) > 8 {
|
||||
data.GemIds = data.GemIds[:8]
|
||||
}
|
||||
user.GemIds = utils.JsonEncode(data.GemIds)
|
||||
}
|
||||
res := h.DB.Updates(&user)
|
||||
if res.Error != nil {
|
||||
resp.ERROR(c, "更新用户信息失败")
|
||||
@@ -584,13 +635,14 @@ func (h *UserHandler) ResetPass(c *gin.Context) {
|
||||
|
||||
session := h.DB.Session(&gorm.Session{})
|
||||
var key string
|
||||
if data.Type == "email" {
|
||||
switch data.Type {
|
||||
case "email":
|
||||
session = session.Where("email", data.Email)
|
||||
key = CodeStorePrefix + data.Email
|
||||
} else if data.Type == "mobile" {
|
||||
case "mobile":
|
||||
session = session.Where("mobile", data.Mobile)
|
||||
key = CodeStorePrefix + data.Mobile
|
||||
} else {
|
||||
default:
|
||||
resp.ERROR(c, "验证类别错误")
|
||||
return
|
||||
}
|
||||
@@ -722,3 +774,81 @@ func (h *UserHandler) SignIn(c *gin.Context) {
|
||||
}
|
||||
resp.SUCCESS(c)
|
||||
}
|
||||
|
||||
// 微信公众号 小程序授权登录
|
||||
func (h *UserHandler) WxAuthLogin(c *gin.Context) {
|
||||
var data struct {
|
||||
Code string `json:"code"`
|
||||
InviteCode string `json:"invite_code"`
|
||||
}
|
||||
if err := c.ShouldBindJSON(&data); err != nil {
|
||||
resp.ERROR(c, types.InvalidArgs)
|
||||
return
|
||||
}
|
||||
|
||||
// 根据 code 获取 openid
|
||||
openID, accessToken, err := h.wechatService.GetOpenIDByCode(data.Code)
|
||||
if err != nil {
|
||||
resp.ERROR(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
// 获取微信用户昵称头像
|
||||
userInfo, err := h.wechatService.GetUserInfo(accessToken, openID)
|
||||
if err != nil {
|
||||
resp.ERROR(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
nickname := userInfo["nickname"].(string)
|
||||
headimgurl := userInfo["headimgurl"].(string)
|
||||
|
||||
// 查询用户是否存在
|
||||
var user model.User
|
||||
h.DB.Where("openid = ?", openID).First(&user)
|
||||
if user.Id > 0 {
|
||||
// 用户存在,更新用户信息并登录
|
||||
user.Nickname = nickname
|
||||
user.Avatar = headimgurl
|
||||
if err := h.DB.Save(&user).Error; err != nil {
|
||||
resp.ERROR(c, "更新用户信息失败")
|
||||
return
|
||||
}
|
||||
|
||||
token, err := h.doLogin(&user, c.ClientIP())
|
||||
if err != nil {
|
||||
resp.ERROR(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
resp.SUCCESS(c, gin.H{"token": token, "user_id": user.Id, "username": user.Username})
|
||||
return
|
||||
}
|
||||
|
||||
// 用户不存在,创建新用户
|
||||
user = model.User{
|
||||
OpenId: openID,
|
||||
Nickname: nickname,
|
||||
Avatar: headimgurl,
|
||||
}
|
||||
|
||||
// 被邀请人也获得赠送算力
|
||||
if data.InviteCode != "" {
|
||||
user.Power = h.App.SysConfig.Base.InitPower * 2
|
||||
}
|
||||
|
||||
user, err = h.createNewUser(user, data.InviteCode)
|
||||
if err != nil {
|
||||
resp.ERROR(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
// 自动登录
|
||||
token, err := h.doLogin(&user, c.ClientIP())
|
||||
if err != nil {
|
||||
resp.ERROR(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
resp.SUCCESS(c, gin.H{"token": token, "user_id": user.Id, "username": user.Username})
|
||||
}
|
||||
|
||||
+127
-155
@@ -20,7 +20,6 @@ import (
|
||||
"geekai/store/vo"
|
||||
"geekai/utils"
|
||||
"geekai/utils/resp"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"gorm.io/gorm"
|
||||
@@ -54,33 +53,50 @@ func (h *VideoHandler) RegisterRoutes() {
|
||||
// 需要用户授权的接口
|
||||
group.Use(middleware.UserAuthMiddleware(h.App.Config.Session.SecretKey, h.App.Redis))
|
||||
{
|
||||
group.POST("luma/create", h.LumaCreate)
|
||||
group.POST("keling/create", h.KeLingCreate)
|
||||
group.POST("create", h.Create)
|
||||
group.GET("list", h.List)
|
||||
group.GET("remove", h.Remove)
|
||||
group.GET("publish", h.Publish)
|
||||
group.GET("power-config", h.GetPowerConfig) // 获取算力配置
|
||||
group.GET("power-by-key", h.GetPowerByPriceKey) // 根据 priceKey 获取算力
|
||||
}
|
||||
}
|
||||
|
||||
func (h *VideoHandler) LumaCreate(c *gin.Context) {
|
||||
type VideoTaskRequest struct {
|
||||
Provider string `json:"provider"` // 服务提供商(不带版本号:veo, sora)
|
||||
Model string `json:"model"` // 模型标识(带版本号:veo-2.0, sora-2.0)
|
||||
Prompt string `json:"prompt"` // 提示词
|
||||
Params map[string]any `json:"params"` // 模型特定参数
|
||||
PriceKey string `json:"price_key"` // 价格键(如 "fixed", "5_720P" 等)
|
||||
}
|
||||
|
||||
var data struct {
|
||||
Prompt string `json:"prompt"`
|
||||
FirstFrameImg string `json:"first_frame_img,omitempty"`
|
||||
EndFrameImg string `json:"end_frame_img,omitempty"`
|
||||
ExpandPrompt bool `json:"expand_prompt,omitempty"`
|
||||
Loop bool `json:"loop,omitempty"`
|
||||
}
|
||||
// Create 统一的创建视频任务接口
|
||||
func (h *VideoHandler) Create(c *gin.Context) {
|
||||
var data VideoTaskRequest
|
||||
if err := c.ShouldBindJSON(&data); err != nil {
|
||||
resp.ERROR(c, types.InvalidArgs)
|
||||
return
|
||||
}
|
||||
// 检查 Prompt 长度
|
||||
|
||||
// 验证必填字段
|
||||
if data.Provider == "" {
|
||||
resp.ERROR(c, "provider 不能为空")
|
||||
return
|
||||
}
|
||||
if data.Model == "" {
|
||||
resp.ERROR(c, "model 不能为空")
|
||||
return
|
||||
}
|
||||
if data.Prompt == "" {
|
||||
resp.ERROR(c, "prompt is needed")
|
||||
resp.ERROR(c, "prompt 不能为空")
|
||||
return
|
||||
}
|
||||
if data.PriceKey == "" {
|
||||
resp.ERROR(c, "price_key 不能为空")
|
||||
return
|
||||
}
|
||||
|
||||
// 文本审查
|
||||
if h.App.SysConfig.Moderation.Enable {
|
||||
moderationResult, err := h.moderationManager.GetService().Moderate(data.Prompt)
|
||||
if err != nil {
|
||||
@@ -101,138 +117,45 @@ func (h *VideoHandler) LumaCreate(c *gin.Context) {
|
||||
resp.ERROR(c, "当前创作内容包含敏感词,请重新输入!")
|
||||
return
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
// 获取用户信息
|
||||
user, err := h.GetLoginUser(c)
|
||||
if err != nil {
|
||||
resp.NotAuth(c)
|
||||
return
|
||||
}
|
||||
|
||||
if user.Power < h.App.SysConfig.Base.LumaPower {
|
||||
resp.ERROR(c, "您的算力不足,请充值后再试!")
|
||||
return
|
||||
}
|
||||
|
||||
userId := int(h.GetLoginUserId(c))
|
||||
params := types.LumaVideoParams{
|
||||
PromptOptimize: data.ExpandPrompt,
|
||||
Loop: data.Loop,
|
||||
StartImgURL: data.FirstFrameImg,
|
||||
EndImgURL: data.EndFrameImg,
|
||||
}
|
||||
task := types.VideoTask{
|
||||
UserId: userId,
|
||||
Type: types.VideoLuma,
|
||||
Prompt: data.Prompt,
|
||||
Params: params,
|
||||
TranslateModelId: h.App.SysConfig.Base.AssistantModelId,
|
||||
}
|
||||
// 插入数据库
|
||||
job := model.VideoJob{
|
||||
UserId: uint(userId),
|
||||
Type: types.VideoLuma,
|
||||
Prompt: data.Prompt,
|
||||
Power: h.App.SysConfig.Base.LumaPower,
|
||||
TaskInfo: utils.JsonEncode(task),
|
||||
}
|
||||
tx := h.DB.Create(&job)
|
||||
if tx.Error != nil {
|
||||
resp.ERROR(c, tx.Error.Error())
|
||||
return
|
||||
}
|
||||
|
||||
// 创建任务
|
||||
task.Id = job.Id
|
||||
h.videoService.PushTask(task)
|
||||
|
||||
// update user's power
|
||||
err = h.userService.DecreasePower(job.UserId, job.Power, model.PowerLog{
|
||||
Type: types.PowerConsume,
|
||||
Model: "luma",
|
||||
Remark: fmt.Sprintf("Luma 文生视频,任务ID:%d", job.Id),
|
||||
})
|
||||
// 计算算力
|
||||
power, err := video.CalculatePower(h.DB, data.Model, data.PriceKey)
|
||||
if err != nil {
|
||||
resp.ERROR(c, err.Error())
|
||||
return
|
||||
}
|
||||
resp.SUCCESS(c)
|
||||
}
|
||||
|
||||
func (h *VideoHandler) KeLingCreate(c *gin.Context) {
|
||||
|
||||
var data struct {
|
||||
Channel string `json:"channel"`
|
||||
TaskType string `json:"task_type"` // 任务类型: text2video/image2video
|
||||
Model string `json:"model"` // 模型: kling-v1-5,kling-v1-6
|
||||
Prompt string `json:"prompt"` // 视频描述
|
||||
NegPrompt string `json:"negative_prompt"` // 负面提示词
|
||||
CfgScale float64 `json:"cfg_scale"` // 相关性系数(0-1)
|
||||
Mode string `json:"mode"` // 生成模式: std/pro
|
||||
AspectRatio string `json:"aspect_ratio"` // 画面比例: 16:9/9:16/1:1
|
||||
Duration string `json:"duration"` // 视频时长: 5/10
|
||||
CameraControl types.CameraControl `json:"camera_control"` // 摄像机控制
|
||||
Image string `json:"image"` // 参考图片URL(image2video)
|
||||
ImageTail string `json:"image_tail"` // 尾帧图片URL(image2video)
|
||||
}
|
||||
if err := c.ShouldBindJSON(&data); err != nil {
|
||||
resp.ERROR(c, types.InvalidArgs)
|
||||
return
|
||||
}
|
||||
|
||||
user, err := h.GetLoginUser(c)
|
||||
if err != nil {
|
||||
resp.NotAuth(c)
|
||||
return
|
||||
}
|
||||
|
||||
// 计算当前任务所需算力
|
||||
key := fmt.Sprintf("%s_%s_%s", data.Model, data.Mode, data.Duration)
|
||||
power := h.App.SysConfig.Base.KeLingPowers[key]
|
||||
if power == 0 {
|
||||
resp.ERROR(c, "当前模型暂不支持")
|
||||
return
|
||||
}
|
||||
// 检查算力是否充足
|
||||
if user.Power < power {
|
||||
resp.ERROR(c, "您的算力不足,请充值后再试!")
|
||||
return
|
||||
}
|
||||
|
||||
if data.Prompt == "" {
|
||||
resp.ERROR(c, "prompt is needed")
|
||||
return
|
||||
}
|
||||
|
||||
// 构建任务
|
||||
userId := int(h.GetLoginUserId(c))
|
||||
params := types.KeLingVideoParams{
|
||||
TaskType: data.TaskType,
|
||||
Model: data.Model,
|
||||
Prompt: data.Prompt,
|
||||
NegPrompt: data.NegPrompt,
|
||||
CfgScale: data.CfgScale,
|
||||
Mode: data.Mode,
|
||||
AspectRatio: data.AspectRatio,
|
||||
Duration: data.Duration,
|
||||
CameraControl: data.CameraControl,
|
||||
Image: data.Image,
|
||||
ImageTail: data.ImageTail,
|
||||
}
|
||||
task := types.VideoTask{
|
||||
UserId: userId,
|
||||
Type: types.VideoKeLing,
|
||||
Type: data.Provider, // provider 作为 type
|
||||
Prompt: data.Prompt,
|
||||
Params: params,
|
||||
Params: data.Params,
|
||||
TranslateModelId: h.App.SysConfig.Base.AssistantModelId,
|
||||
Channel: data.Channel,
|
||||
}
|
||||
|
||||
// 插入数据库
|
||||
job := model.VideoJob{
|
||||
UserId: uint(userId),
|
||||
Type: types.VideoKeLing,
|
||||
Prompt: data.Prompt,
|
||||
Power: power,
|
||||
TaskInfo: utils.JsonEncode(task),
|
||||
UserId: uint(userId),
|
||||
Type: data.Provider,
|
||||
Prompt: data.Prompt,
|
||||
Power: power,
|
||||
Params: utils.JsonEncode(task),
|
||||
}
|
||||
tx := h.DB.Create(&job)
|
||||
if tx.Error != nil {
|
||||
@@ -244,17 +167,52 @@ func (h *VideoHandler) KeLingCreate(c *gin.Context) {
|
||||
task.Id = job.Id
|
||||
h.videoService.PushTask(task)
|
||||
|
||||
// update user's power
|
||||
// 扣减算力
|
||||
err = h.userService.DecreasePower(job.UserId, job.Power, model.PowerLog{
|
||||
Type: types.PowerConsume,
|
||||
Model: "keling",
|
||||
Remark: fmt.Sprintf("keling 文生视频,任务ID:%d", job.Id),
|
||||
Model: data.Provider,
|
||||
Remark: fmt.Sprintf("%s 视频生成,任务ID:%d", data.Provider, job.Id),
|
||||
})
|
||||
if err != nil {
|
||||
resp.ERROR(c, err.Error())
|
||||
return
|
||||
}
|
||||
resp.SUCCESS(c)
|
||||
|
||||
resp.SUCCESS(c, gin.H{"job_id": job.Id})
|
||||
}
|
||||
|
||||
// GetPowerConfig 获取算力配置
|
||||
func (h *VideoHandler) GetPowerConfig(c *gin.Context) {
|
||||
config, err := video.GetVideoConfig(h.DB)
|
||||
if err != nil {
|
||||
resp.ERROR(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
resp.SUCCESS(c, config.VideoPowers)
|
||||
}
|
||||
|
||||
// GetPowerByPriceKey 根据 modelKey 和 priceKey 获取算力值
|
||||
func (h *VideoHandler) GetPowerByPriceKey(c *gin.Context) {
|
||||
modelKey := c.Query("model_key")
|
||||
priceKey := c.Query("price_key")
|
||||
|
||||
if modelKey == "" {
|
||||
resp.ERROR(c, "model_key 不能为空")
|
||||
return
|
||||
}
|
||||
if priceKey == "" {
|
||||
resp.ERROR(c, "price_key 不能为空")
|
||||
return
|
||||
}
|
||||
|
||||
power, err := video.CalculatePower(h.DB, modelKey, priceKey)
|
||||
if err != nil {
|
||||
resp.ERROR(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
resp.SUCCESS(c, gin.H{"power": power})
|
||||
}
|
||||
|
||||
func (h *VideoHandler) List(c *gin.Context) {
|
||||
@@ -268,7 +226,7 @@ func (h *VideoHandler) List(c *gin.Context) {
|
||||
session = session.Where("type", t)
|
||||
}
|
||||
if all {
|
||||
session = session.Where("publish", 0).Where("progress", 100)
|
||||
session = session.Where("publish", 0).Where("status", types.VideoStatusSuccess)
|
||||
} else {
|
||||
session = session.Where("user_id", userId)
|
||||
}
|
||||
@@ -296,36 +254,51 @@ func (h *VideoHandler) List(c *gin.Context) {
|
||||
continue
|
||||
}
|
||||
item.CreatedAt = v.CreatedAt.Unix()
|
||||
if item.VideoURL == "" {
|
||||
item.VideoURL = v.WaterURL
|
||||
}
|
||||
// 解析任务详情
|
||||
if item.Type == types.VideoKeLing {
|
||||
// 解析任务详情(用于前端展示标签)
|
||||
if v.Params != "" {
|
||||
task := types.VideoTask{}
|
||||
err = utils.JsonDecode(v.TaskInfo, &task)
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
var params types.KeLingVideoParams
|
||||
err = utils.JsonDecode(utils.JsonEncode(task.Params), ¶ms)
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
item.RawData = map[string]interface{}{
|
||||
"task_type": params.TaskType,
|
||||
"model": params.Model,
|
||||
"cfg_scale": params.CfgScale,
|
||||
"mode": params.Mode,
|
||||
"aspect_ratio": params.AspectRatio,
|
||||
"duration": params.Duration,
|
||||
"model_name": fmt.Sprintf("%s_%s_%s", params.Model, params.Mode, params.Duration),
|
||||
}
|
||||
|
||||
// 如果视频URL不为空,则设置为生成成功
|
||||
if item.VideoURL != "" {
|
||||
item.Progress = 100
|
||||
if err := utils.JsonDecode(v.Params, &task); err == nil {
|
||||
// 默认从 params map 中提取常用字段
|
||||
if paramsMap, ok := task.Params.(map[string]any); ok {
|
||||
if item.Params == nil {
|
||||
item.Params = make(map[string]any)
|
||||
}
|
||||
if _, ok := item.Params["task_type"]; !ok {
|
||||
if taskType, ok := paramsMap["task_type"]; ok {
|
||||
item.Params["task_type"] = taskType
|
||||
}
|
||||
}
|
||||
if _, ok := item.Params["model"]; !ok {
|
||||
if modelKey, ok := paramsMap["model"]; ok {
|
||||
item.Params["model"] = modelKey
|
||||
}
|
||||
}
|
||||
if _, ok := item.Params["duration"]; !ok {
|
||||
if duration, ok := paramsMap["duration"]; ok {
|
||||
item.Params["duration"] = duration
|
||||
}
|
||||
}
|
||||
if _, ok := item.Params["size"]; !ok {
|
||||
if size, ok := paramsMap["size"]; ok {
|
||||
item.Params["size"] = size
|
||||
} else if size, ok := paramsMap["resolution"].(string); ok {
|
||||
item.Params["size"] = size
|
||||
}
|
||||
}
|
||||
if _, ok := item.Params["mode"]; !ok {
|
||||
if mode, ok := paramsMap["mode"]; ok {
|
||||
item.Params["mode"] = mode
|
||||
}
|
||||
}
|
||||
if _, ok := item.Params["sound"]; !ok {
|
||||
if sound, ok := paramsMap["sound"]; ok {
|
||||
item.Params["sound"] = sound
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
items = append(items, item)
|
||||
}
|
||||
|
||||
@@ -341,9 +314,9 @@ func (h *VideoHandler) Remove(c *gin.Context) {
|
||||
resp.ERROR(c, err.Error())
|
||||
return
|
||||
}
|
||||
// 只有失败或者超时的任务才能删除
|
||||
if !(job.Progress == service.FailTaskProgress || time.Now().After(job.CreatedAt.Add(time.Minute*30))) {
|
||||
resp.ERROR(c, "只有失败和超时(30分钟)的任务才能删除!")
|
||||
// 只有失败的任务才能删除
|
||||
if job.Status != types.VideoStatusFailed {
|
||||
resp.ERROR(c, "只有失败的任务才能删除!")
|
||||
return
|
||||
}
|
||||
|
||||
@@ -355,7 +328,6 @@ func (h *VideoHandler) Remove(c *gin.Context) {
|
||||
}
|
||||
|
||||
// 删除文件
|
||||
_ = h.uploader.GetUploadHandler().Delete(job.CoverURL)
|
||||
_ = h.uploader.GetUploadHandler().Delete(job.VideoURL)
|
||||
|
||||
resp.SUCCESS(c)
|
||||
|
||||
@@ -0,0 +1,95 @@
|
||||
package handler
|
||||
|
||||
// * +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
|
||||
// * Copyright 2023 The Geek-AI Authors. All rights reserved.
|
||||
// * Use of this source code is governed by a Apache-2.0 license
|
||||
// * that can be found in the LICENSE file.
|
||||
// * @Author yangjian102621@163.com
|
||||
// * +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
|
||||
|
||||
import (
|
||||
"crypto/sha1"
|
||||
"encoding/hex"
|
||||
"geekai/core"
|
||||
"geekai/utils/resp"
|
||||
"log"
|
||||
"net/http"
|
||||
"sort"
|
||||
"strings"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
type WxGzhHandler struct {
|
||||
BaseHandler
|
||||
}
|
||||
|
||||
func NewWxGzhHandler(server *core.AppServer) *WxGzhHandler {
|
||||
return &WxGzhHandler{
|
||||
BaseHandler: BaseHandler{
|
||||
App: server,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func (h *WxGzhHandler) RegisterRoutes() {
|
||||
group := h.App.Engine.Group("/api/wx/")
|
||||
group.GET("verify", h.WechatVerify)
|
||||
}
|
||||
|
||||
// 处理微信服务器验证请求
|
||||
func (h *WxGzhHandler) WechatVerify(c *gin.Context) {
|
||||
logger.Info("WechatVerify")
|
||||
// 只处理 GET 请求
|
||||
if c.Request.Method != "GET" {
|
||||
resp.ERROR(c, "Method Not Allowed")
|
||||
return
|
||||
}
|
||||
|
||||
// 解析 URL 参数
|
||||
signature := c.Query("signature")
|
||||
timestamp := c.Query("timestamp")
|
||||
nonce := c.Query("nonce")
|
||||
echostr := c.Query("echostr")
|
||||
|
||||
// 验证参数完整性
|
||||
if signature == "" || timestamp == "" || nonce == "" || echostr == "" {
|
||||
log.Println("Missing parameters")
|
||||
resp.ERROR(c, "Missing parameters")
|
||||
return
|
||||
}
|
||||
|
||||
// 验证签名
|
||||
if validateSignature(signature, h.App.SysConfig.WxGzh.Token, timestamp, nonce) {
|
||||
// 验证成功,返回 echostr(必须是纯文本)
|
||||
c.String(http.StatusOK, echostr)
|
||||
log.Println("Token verification success")
|
||||
} else {
|
||||
// 验证失败
|
||||
resp.ERROR(c, "Forbidden: Invalid signature")
|
||||
log.Println("Token verification failed")
|
||||
}
|
||||
}
|
||||
|
||||
func validateSignature(signature, token, timestamp, nonce string) bool {
|
||||
// 1. 将 token、timestamp、nonce 按字典序排序
|
||||
strs := []string{token, timestamp, nonce}
|
||||
sort.Strings(strs)
|
||||
|
||||
// 2. 拼接字符串
|
||||
joined := strings.Join(strs, "")
|
||||
|
||||
// 3. 计算 SHA1 哈希
|
||||
hash := sha1.New()
|
||||
hash.Write([]byte(joined))
|
||||
hashed := hex.EncodeToString(hash.Sum(nil))
|
||||
|
||||
// 4. 与 signature 比对
|
||||
return hashed == signature
|
||||
}
|
||||
|
||||
// 创建微信菜单
|
||||
func (h *WxGzhHandler) CreateMenu(c *gin.Context) {
|
||||
|
||||
resp.SUCCESS(c, "创建菜单成功")
|
||||
}
|
||||
Reference in New Issue
Block a user