mirror of
https://github.com/yangjian102621/geekai.git
synced 2026-08-14 03:31:00 +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)
|
||||
}
|
||||
Reference in New Issue
Block a user