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:
RockYang
2026-08-11 14:51:20 +08:00
parent 37e024acea
commit 9ccff4efbc
219 changed files with 29212 additions and 10802 deletions
@@ -1,4 +1,4 @@
package dalle
package image
// * +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
// * Copyright 2023 The Geek-AI Authors. All rights reserved.
@@ -10,12 +10,13 @@ package dalle
import (
"fmt"
"geekai/core/types"
logger2 "geekai/logger"
"geekai/log"
"geekai/service"
"geekai/service/oss"
"geekai/store"
"geekai/store/model"
"geekai/utils"
"io"
"time"
"github.com/go-redis/redis/v8"
@@ -24,9 +25,9 @@ import (
"gorm.io/gorm"
)
var logger = logger2.GetLogger()
var logger = log.GetLogger()
// DALL-E 绘画服务
// Image Generation Service
type Service struct {
httpClient *req.Client
@@ -40,50 +41,50 @@ func NewService(db *gorm.DB, manager *oss.UploaderManager, redisCli *redis.Clien
return &Service{
httpClient: req.C().SetTimeout(time.Minute * 3),
db: db,
taskQueue: store.NewRedisQueue("DallE_Task_Queue", redisCli),
taskQueue: store.NewRedisQueue("Image_Task_Queue", redisCli),
uploadManager: manager,
userService: userService,
}
}
// PushTask push a new mj task in to task queue
func (s *Service) PushTask(task types.DallTask) {
logger.Infof("add a new DALL-E task to the task list: %+v", task)
// PushTask push a new image task in to task queue
func (s *Service) PushTask(task types.ImageTask) {
logger.Infof("add a new Image generation task to the task list: %+v", task)
if err := s.taskQueue.RPush(task); err != nil {
logger.Errorf("push dall-e task to queue failed: %v", err)
logger.Errorf("push image task to queue failed: %v", err)
}
}
func (s *Service) Run() {
// 将数据库中未提交的任务加载到队列
var jobs []model.DallJob
var jobs []model.ImageJob
s.db.Where("progress", 0).Find(&jobs)
for _, v := range jobs {
var task types.DallTask
err := utils.JsonDecode(v.TaskInfo, &task)
var task types.ImageTask
err := utils.JsonDecode(v.Params, &task)
if err != nil {
logger.Errorf("decode task info with error: %v", err)
logger.Errorf("decode task params with error: %v", err)
continue
}
task.Id = v.Id
s.PushTask(task)
}
logger.Info("Starting DALL-E job consumer...")
logger.Info("Starting Image generation job consumer...")
go func() {
for {
var task types.DallTask
var task types.ImageTask
err := s.taskQueue.LPop(&task)
if err != nil {
logger.Errorf("taking task with error: %v", err)
continue
}
logger.Infof("handle a new DALL-E task: %+v", task)
logger.Infof("handle a new Image generation task: %+v", task)
go func() {
_, err = s.Image(task, false)
if err != nil {
logger.Errorf("error with image task: %v", err)
s.db.Model(&model.DallJob{Id: task.Id}).UpdateColumns(map[string]interface{}{
s.db.Model(&model.ImageJob{Id: task.Id}).UpdateColumns(map[string]interface{}{
"progress": service.FailTaskProgress,
"err_msg": err.Error(),
})
@@ -120,7 +121,7 @@ type ErrRes struct {
} `json:"error"`
}
func (s *Service) Image(task types.DallTask, sync bool) (string, error) {
func (s *Service) Image(task types.ImageTask, sync bool) (string, error) {
logger.Debugf("绘画参数:%+v", task)
var chatModel model.ChatModel
@@ -136,7 +137,7 @@ func (s *Service) Image(task types.DallTask, sync bool) (string, error) {
if chatModel.KeyId > 0 {
session = session.Where("id = ?", chatModel.KeyId)
} else {
session = session.Where("type = ?", "dalle")
session = session.Where("type = ?", "image")
}
err := session.Order("last_used_at ASC").First(&apiKey).Error
if err != nil {
@@ -179,6 +180,11 @@ func (s *Service) Image(task types.DallTask, sync bool) (string, error) {
return "", fmt.Errorf("error with send request, status: %s, %+v", r.Status, errRes.Error)
}
if len(res.Data) == 0 && r.Body != nil {
body, _ := io.ReadAll(r.Body)
return "", fmt.Errorf("%s", string(body))
}
// update the api key last use time
s.db.Model(&apiKey).UpdateColumn("last_used_at", time.Now().Unix())
var imgURL string
@@ -199,7 +205,7 @@ func (s *Service) Image(task types.DallTask, sync bool) (string, error) {
}
data["org_url"] = imgURL
// update task progress
err = s.db.Model(&model.DallJob{Id: task.Id}).UpdateColumns(data).Error
err = s.db.Model(&model.ImageJob{Id: task.Id}).UpdateColumns(data).Error
if err != nil {
return "", fmt.Errorf("err with update database: %v", err)
}
@@ -214,10 +220,10 @@ func (s *Service) Image(task types.DallTask, sync bool) (string, error) {
func (s *Service) CheckTaskStatus() {
go func() {
logger.Info("Running DALL-E task status checking ...")
logger.Info("Running Image generation task status checking ...")
for {
// 检查未完成任务进度
var jobs []model.DallJob
var jobs []model.ImageJob
s.db.Where("progress < ?", 100).Find(&jobs)
for _, job := range jobs {
// 超时的任务标记为失败
@@ -231,8 +237,8 @@ func (s *Service) CheckTaskStatus() {
// 找出失败的任务,并恢复其扣减算力
s.db.Where("progress", service.FailTaskProgress).Where("power > ?", 0).Find(&jobs)
for _, job := range jobs {
var task types.DallTask
err := utils.JsonDecode(job.TaskInfo, &task)
var task types.ImageTask
err := utils.JsonDecode(job.Params, &task)
if err != nil {
continue
}
@@ -254,7 +260,7 @@ func (s *Service) CheckTaskStatus() {
func (s *Service) DownloadImages() {
go func() {
var items []model.DallJob
var items []model.ImageJob
for {
res := s.db.Where("img_url = ? AND progress = ?", "", 100).Find(&items)
if res.Error != nil {
@@ -291,7 +297,7 @@ func (s *Service) downloadImage(jobId uint, orgURL string) (string, error) {
}
// update img_url
res := s.db.Model(&model.DallJob{Id: jobId}).UpdateColumn("img_url", imgURL)
res := s.db.Model(&model.ImageJob{Id: jobId}).UpdateColumn("img_url", imgURL)
if res.Error != nil {
return "", err
}
@@ -10,7 +10,7 @@ import (
"gorm.io/gorm"
"geekai/core/types"
logger2 "geekai/logger"
"geekai/log"
"geekai/service"
"geekai/service/oss"
"geekai/store"
@@ -20,7 +20,7 @@ import (
"github.com/go-redis/redis/v8"
)
var logger = logger2.GetLogger()
var logger = log.GetLogger()
// Service 即梦服务(合并了消费者功能)
type Service struct {
@@ -249,8 +249,10 @@ func (s *Service) buildTaskRequest(req *types.JimengTaskRequest) (map[string]any
// duration 转成 frames
if duration, ok := params["duration"]; ok {
if secs, ok := duration.(int); ok {
params["frames"] = secs*24 + 1
if v, ok := duration.(int); ok {
params["frames"] = v*24 + 1
} else if v, ok := duration.(float64); ok {
params["frames"] = int(v*24) + 1
}
delete(params, "duration")
}
+420 -123
View File
@@ -8,17 +8,15 @@ package service
// ++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
import (
"context"
"encoding/json"
"fmt"
"geekai/core/types"
"geekai/store"
"geekai/store/model"
"strings"
"sync"
"github.com/go-redis/redis/v8"
"gorm.io/gorm"
"gorm.io/gorm/schema"
)
const (
@@ -33,23 +31,21 @@ type MigrationService struct {
db *gorm.DB
redisClient *redis.Client
appConfig *types.AppConfig
levelDB *store.LevelDB
}
func NewMigrationService(db *gorm.DB, redisClient *redis.Client, appConfig *types.AppConfig, levelDB *store.LevelDB) *MigrationService {
func NewMigrationService(db *gorm.DB, redisClient *redis.Client, appConfig *types.AppConfig) *MigrationService {
return &MigrationService{
db: db,
redisClient: redisClient,
appConfig: appConfig,
levelDB: levelDB,
}
}
func (s *MigrationService) StartMigrate() {
// 表结构同步必须在对外服务前完成,避免缺列导致业务报错
// 表结构迁移必须在业务服务启动前完成,避免新表和新列尚未创建就被后台任务查询。
s.TableMigration()
go func() {
_ = s.MigrateConfig(s.appConfig)
s.MigrateConfig(s.appConfig)
}()
}
@@ -126,163 +122,386 @@ func (s *MigrationService) MigrateConfigContent() error {
return nil
}
// 永不删除的保护列(大小写不敏感
var protectedColumns = map[string]struct{}{
"id": {},
"created_at": {},
"updated_at": {},
}
// allModels 全部需要同步的数据表 model
func allModels() []any {
return []any{
// fullTableMigration 第一步:全量表迁移,同步所有表结构(新增表、新增字段、字段类型
// 适用于首次安装或导入旧版数据库后同步 schema,AutoMigrate 会补齐缺失的表和列
func (s *MigrationService) fullTableMigration() {
logger.Info("执行全量表迁移(同步 schema...")
models := []any{
&model.Config{},
&model.AdminUser{},
&model.ChatApp{},
&model.ApiKey{},
&model.AppType{},
&model.ChatApp{},
&model.ChatModel{},
&model.User{},
&model.ChatItem{},
&model.ChatMessage{},
&model.ChatModel{},
&model.Config{},
&model.DallJob{},
&model.File{},
&model.Order{},
&model.Product{},
&model.Function{},
&model.Menu{},
&model.InviteCode{},
&model.InviteLog{},
&model.JimengJob{},
&model.Menu{},
&model.MidJourneyJob{},
&model.Moderation{},
&model.Order{},
&model.PowerLog{},
&model.Product{},
&model.Redeem{},
&model.SdJob{},
&model.SunoJob{},
&model.User{},
&model.PowerLog{},
&model.File{},
&model.UserLoginLog{},
&model.MidJourneyJob{},
&model.SunoJob{},
&model.VideoJob{},
&model.JimengJob{},
&model.PPTJob{},
&model.Moderation{},
&model.ImageJob{},
}
}
// 数据表迁移:先处理字段重命名(保数据),再全量同步 schema
func (s *MigrationService) TableMigration() {
logger.Info("开始数据表迁移...")
s.renameColumns()
if err := s.SyncAllModels(); err != nil {
logger.Errorf("同步数据表字段失败: %v", err)
if err := s.db.AutoMigrate(models...); err != nil {
logger.Errorf("全量表迁移失败: %v", err)
return
}
logger.Info("数据表迁移完成")
logger.Info("全量表迁移完成")
}
// renameColumns 只处理「改名」场景:删旧加新会丢数据,必须先 Rename
func (s *MigrationService) renameColumns() {
m := s.db.Migrator()
// fixTableConstraints 第一步之后:根据模型定义修复各表的主键、自增和关键索引
// 主要解决初始化 SQL 中缺少 AUTO_INCREMENT 或 PRIMARY KEY 导致插入失败的问题
func (s *MigrationService) fixTableConstraints() {
logger.Info("开始修复各表的主键、自增属性和索引...")
if m.HasColumn(&model.JimengJob{}, "task_params") {
_ = m.RenameColumn(&model.JimengJob{}, "task_params", "params")
// 当前数据库名
var dbName string
if err := s.db.Raw("SELECT DATABASE()").Scan(&dbName).Error; err != nil {
logger.Errorf("获取当前数据库名失败: %v", err)
return
}
if m.HasColumn(&model.Order{}, "pay_type") {
_ = m.RenameColumn(&model.Order{}, "pay_type", "channel")
if dbName == "" {
logger.Warn("当前连接未选择数据库,跳过约束修复")
return
}
if m.HasColumn(&model.Config{}, "config_json") {
_ = m.RenameColumn(&model.Config{}, "config_json", "value")
}
if m.HasColumn(&model.Config{}, "marker") {
_ = m.RenameColumn(&model.Config{}, "marker", "name")
}
if m.HasIndex(&model.Config{}, "idx_chatgpt_configs_key") {
_ = m.DropIndex(&model.Config{}, "idx_chatgpt_configs_key")
}
if m.HasIndex(&model.Config{}, "marker") {
_ = m.DropIndex(&model.Config{}, "marker")
}
}
// SyncAllModels 按 model 定义同步所有数据表:缺列新建,多余列删除
func (s *MigrationService) SyncAllModels() error {
var firstErr error
for _, m := range allModels() {
if err := s.syncModel(m); err != nil {
logger.Errorf("同步 model %T 失败: %v", m, err)
if firstErr == nil {
firstErr = err
// === 修复所有包含 id 字段的表的主键 + 自增 ===
type columnInfo struct {
TableName string `gorm:"column:TABLE_NAME"`
ColumnName string `gorm:"column:COLUMN_NAME"`
ColumnKey string `gorm:"column:COLUMN_KEY"`
Extra string `gorm:"column:EXTRA"`
DataType string `gorm:"column:DATA_TYPE"`
}
var idColumns []columnInfo
if err := s.db.Raw(`
SELECT TABLE_NAME, COLUMN_NAME, COLUMN_KEY, EXTRA, DATA_TYPE
FROM INFORMATION_SCHEMA.COLUMNS
WHERE TABLE_SCHEMA = ? AND COLUMN_NAME = 'id'
`, dbName).Scan(&idColumns).Error; err != nil {
logger.Errorf("查询各表 id 字段信息失败: %v", err)
} else {
for _, col := range idColumns {
if col.ColumnName == "" {
continue
}
// 检查当前表是否已经存在主键
type pkInfo struct {
ColumnName string `gorm:"column:COLUMN_NAME"`
}
var pkColumns []pkInfo
if err := s.db.Raw(`
SELECT COLUMN_NAME
FROM INFORMATION_SCHEMA.KEY_COLUMN_USAGE
WHERE TABLE_SCHEMA = ? AND TABLE_NAME = ? AND CONSTRAINT_NAME = 'PRIMARY'
`, dbName, col.TableName).Scan(&pkColumns).Error; err != nil {
logger.Errorf("查询表 %s 的主键信息失败: %v", col.TableName, err)
continue
}
hasPK := len(pkColumns) > 0
pkOnIdOnly := hasPK && len(pkColumns) == 1 && pkColumns[0].ColumnName == "id"
needAlter := false
switch {
case !hasPK:
// 没有任何主键:允许将 id 设置为自增主键
if col.ColumnKey != "PRI" || col.Extra == "" || !strings.Contains(col.Extra, "auto_increment") {
needAlter = true
}
case pkOnIdOnly:
// 只有 id 作为主键:只补充自增属性
if col.Extra == "" || !strings.Contains(col.Extra, "auto_increment") {
needAlter = true
}
default:
// 已存在非 id 或组合主键:避免破坏原有主键,直接跳过
logger.Infof("表 %s 已存在非 id 主键,跳过 id 自增主键修复", col.TableName)
}
if !needAlter {
continue
}
// 维持原来的数据类型,避免与历史 SQL 冲突
dataType := col.DataType
if dataType == "" {
dataType = "int"
}
alterSQL := fmt.Sprintf(
"ALTER TABLE `%s` MODIFY COLUMN id %s NOT NULL AUTO_INCREMENT PRIMARY KEY",
col.TableName,
dataType,
)
if err := s.db.Exec(alterSQL).Error; err != nil {
logger.Errorf("修复表 %s 的 id 自增主键失败: %v", col.TableName, err)
} else {
logger.Infof("已修复表 %s 的 id 为 AUTO_INCREMENT PRIMARY KEY", col.TableName)
}
}
}
return firstErr
// === 为 geekai_users.username 补充唯一索引(根据模型 uniqueIndex 定义) ===
var usernameUniqueCount int64
err := s.db.Raw(`
SELECT COUNT(1)
FROM INFORMATION_SCHEMA.STATISTICS
WHERE TABLE_SCHEMA = ?
AND TABLE_NAME = 'geekai_users'
AND COLUMN_NAME = 'username'
AND NON_UNIQUE = 0
`, dbName).Scan(&usernameUniqueCount).Error
if err != nil {
logger.Errorf("检查 geekai_users.username 唯一索引失败: %v", err)
} else if usernameUniqueCount == 0 {
// 索引名尽量固定,避免重复创建
if err := s.db.Exec("ALTER TABLE geekai_users ADD UNIQUE KEY idx_geekai_users_username (username)").Error; err != nil {
logger.Errorf("创建 geekai_users.username 唯一索引失败: %v", err)
} else {
logger.Info("已为 geekai_users.username 创建唯一索引 idx_geekai_users_username")
}
} else {
logger.Info("geekai_users.username 唯一索引已存在,跳过创建")
}
logger.Info("关键表主键、自增和索引修复完成")
}
func (s *MigrationService) syncModel(dst any) error {
tableName := s.tableName(dst)
if err := s.db.AutoMigrate(dst); err != nil {
return fmt.Errorf("AutoMigrate %s: %w", tableName, err)
// incrementalTableMigration 第二步:增量迁移,仅处理删除字段与数据迁移
// AutoMigrate 不会删除列,故需在此显式 DropColumn;字段重命名需先拷贝数据再删除旧列
func (s *MigrationService) incrementalTableMigration() {
logger.Info("执行增量迁移(删除字段 + 数据迁移)...")
// ========== 字段重命名:全量迁移已添加新列,需将旧列数据拷贝到新列后删除旧列 ==========
// ChatApp(geekai_chat_roles): context_json -> system_prompt 历史数据迁移
if s.db.Migrator().HasColumn(&model.ChatApp{}, "context_json") {
// 将旧列 context_json 的值拷贝到 system_promptNULL 转为空字符串,保证 NOT NULL 约束)
if err := s.db.Exec(`
UPDATE geekai_chat_roles
SET system_prompt = IFNULL(NULLIF(TRIM(COALESCE(context_json, '')), ''), '')
`).Error; err != nil {
logger.Errorf("迁移 geekai_chat_roles.context_json -> system_prompt 失败: %v", err)
} else {
if err := s.db.Migrator().DropColumn(&model.ChatApp{}, "context_json"); err != nil {
logger.Errorf("删除 geekai_chat_roles.context_json 失败: %v", err)
} else {
logger.Info("geekai_chat_roles: context_json 已迁移至 system_prompt 并删除旧列")
}
}
}
if err := s.dropUnusedColumns(dst); err != nil {
return fmt.Errorf("drop unused columns %s: %w", tableName, err)
// ChatApp: 将 user_id 为 NULL 的历史记录置为 0(系统内置)
if s.db.Migrator().HasColumn(&model.ChatApp{}, "user_id") {
if err := s.db.Exec(`UPDATE geekai_chat_roles SET user_id = 0 WHERE user_id IS NULL`).Error; err != nil {
logger.Errorf("初始化 geekai_chat_roles.user_id 失败: %v", err)
}
}
logger.Infof("已同步数据表: %s", tableName)
return nil
// ChatApp(geekai_chat_roles): 删除 marker 列(应用仅通过 id 区分)
var hasMarker int
if s.db.Raw("SELECT COUNT(1) FROM INFORMATION_SCHEMA.COLUMNS WHERE TABLE_SCHEMA = DATABASE() AND TABLE_NAME = 'geekai_chat_roles' AND COLUMN_NAME = 'marker'").Scan(&hasMarker).Error == nil && hasMarker > 0 {
// 先删除可能存在的唯一索引(不同版本 SQL 索引名不同)
for _, idxName := range []string{"marker", "idx_chatgpt_chat_roles_marker", "idx_chatgpt_chat_roles_key", "idx_geekai_chat_roles_marker"} {
_ = s.db.Exec(fmt.Sprintf("ALTER TABLE geekai_chat_roles DROP INDEX `%s`", idxName)).Error
}
if err := s.db.Exec("ALTER TABLE geekai_chat_roles DROP COLUMN marker").Error; err != nil {
logger.Errorf("删除 geekai_chat_roles.marker 失败: %v", err)
} else {
logger.Info("geekai_chat_roles: 已删除 marker 列")
}
}
// Config: config_json -> value, marker -> name
if s.db.Migrator().HasColumn(&model.Config{}, "config_json") {
s.db.Exec("UPDATE geekai_configs SET `value` = config_json WHERE config_json IS NOT NULL AND config_json != ''")
s.db.Migrator().DropColumn(&model.Config{}, "config_json")
}
if s.db.Migrator().HasColumn(&model.Config{}, "marker") {
s.db.Exec("UPDATE geekai_configs SET `name` = marker WHERE marker IS NOT NULL AND marker != ''")
s.db.Migrator().DropColumn(&model.Config{}, "marker")
}
if s.db.Migrator().HasIndex(&model.Config{}, "idx_chatgpt_configs_key") {
s.db.Migrator().DropIndex(&model.Config{}, "idx_chatgpt_configs_key")
}
if s.db.Migrator().HasIndex(&model.Config{}, "marker") {
s.db.Migrator().DropIndex(&model.Config{}, "marker")
}
// Order: pay_type -> channel
if s.db.Migrator().HasColumn(&model.Order{}, "pay_type") {
s.db.Exec("UPDATE geekai_orders SET channel = pay_type WHERE pay_type IS NOT NULL AND pay_type != ''")
s.db.Migrator().DropColumn(&model.Order{}, "pay_type")
}
// JimengJob: task_params -> params
if s.db.Migrator().HasColumn(&model.JimengJob{}, "task_params") {
s.db.Exec("UPDATE geekai_jimeng_jobs SET params = task_params WHERE task_params IS NOT NULL AND task_params != ''")
s.db.Migrator().DropColumn(&model.JimengJob{}, "task_params")
}
// VideoJob: task_info -> params, raw_data -> output
if s.db.Migrator().HasColumn(&model.VideoJob{}, "task_info") {
s.db.Exec("UPDATE geekai_video_jobs SET params = task_info WHERE task_info IS NOT NULL AND task_info != ''")
s.db.Migrator().DropColumn(&model.VideoJob{}, "task_info")
}
if s.db.Migrator().HasColumn(&model.VideoJob{}, "raw_data") {
s.db.Exec("UPDATE geekai_video_jobs SET `output` = raw_data WHERE raw_data IS NOT NULL AND raw_data != ''")
s.db.Migrator().DropColumn(&model.VideoJob{}, "raw_data")
}
// SunoJob: task_info -> params, raw_data -> output
if s.db.Migrator().HasColumn(&model.SunoJob{}, "task_info") {
s.db.Exec("UPDATE geekai_suno_jobs SET params = task_info WHERE task_info IS NOT NULL AND task_info != ''")
s.db.Migrator().DropColumn(&model.SunoJob{}, "task_info")
}
if s.db.Migrator().HasColumn(&model.SunoJob{}, "raw_data") {
s.db.Exec("UPDATE geekai_suno_jobs SET `output` = raw_data WHERE raw_data IS NOT NULL AND raw_data != ''")
s.db.Migrator().DropColumn(&model.SunoJob{}, "raw_data")
}
// ========== 删除不再使用的字段 ==========
if s.db.Migrator().HasColumn(&model.Order{}, "deleted_at") {
s.db.Migrator().DropColumn(&model.Order{}, "deleted_at")
}
if s.db.Migrator().HasColumn(&model.ChatItem{}, "deleted_at") {
s.db.Migrator().DropColumn(&model.ChatItem{}, "deleted_at")
}
if s.db.Migrator().HasColumn(&model.ChatMessage{}, "deleted_at") {
s.db.Migrator().DropColumn(&model.ChatMessage{}, "deleted_at")
}
if s.db.Migrator().HasColumn(&model.User{}, "chat_config") {
s.db.Migrator().DropColumn(&model.User{}, "chat_config")
}
if s.db.Migrator().HasColumn(&model.ChatModel{}, "category") {
s.db.Migrator().DropColumn(&model.ChatModel{}, "category")
}
if s.db.Migrator().HasColumn(&model.ChatModel{}, "description") {
s.db.Migrator().DropColumn(&model.ChatModel{}, "description")
}
if s.db.Migrator().HasColumn(&model.Product{}, "discount") {
s.db.Migrator().DropColumn(&model.Product{}, "discount")
}
if s.db.Migrator().HasColumn(&model.Product{}, "days") {
s.db.Migrator().DropColumn(&model.Product{}, "days")
}
if s.db.Migrator().HasColumn(&model.Product{}, "app_url") {
s.db.Migrator().DropColumn(&model.Product{}, "app_url")
}
if s.db.Migrator().HasColumn(&model.Product{}, "url") {
s.db.Migrator().DropColumn(&model.Product{}, "url")
}
if s.db.Migrator().HasColumn(&model.VideoJob{}, "water_url") {
s.db.Migrator().DropColumn(&model.VideoJob{}, "water_url")
}
if s.db.Migrator().HasColumn(&model.VideoJob{}, "cover_url") {
s.db.Migrator().DropColumn(&model.VideoJob{}, "cover_url")
}
if s.db.Migrator().HasColumn(&model.VideoJob{}, "prompt_ext") {
s.db.Migrator().DropColumn(&model.VideoJob{}, "prompt_ext")
}
if s.db.Migrator().HasColumn(&model.SunoJob{}, "instrumental") {
s.db.Migrator().DropColumn(&model.SunoJob{}, "instrumental")
}
if s.db.Migrator().HasColumn(&model.SunoJob{}, "tags") {
s.db.Migrator().DropColumn(&model.SunoJob{}, "tags")
}
if s.db.Migrator().HasColumn(&model.SunoJob{}, "extend_secs") {
s.db.Migrator().DropColumn(&model.SunoJob{}, "extend_secs")
}
if s.db.Migrator().HasColumn(&model.SunoJob{}, "model_name") {
s.db.Migrator().DropColumn(&model.SunoJob{}, "model_name")
}
// ========== 数据迁移:根据业务逻辑更新现有数据 ==========
// video_job: 根据 progress 填充 status
if s.db.Migrator().HasColumn(&model.VideoJob{}, "status") {
s.db.Exec(`UPDATE geekai_video_jobs SET status = CASE
WHEN progress < 100 THEN 'in_progress'
WHEN progress = 100 THEN 'success'
WHEN progress = 101 THEN 'failed'
WHEN progress = 102 THEN 'downloading'
ELSE 'pending'
END WHERE status = '' OR status IS NULL`)
}
// suno_job: 从 output 提取 tags/model_name 填入 params
s.migrateSunoJobData()
logger.Info("增量迁移完成")
}
func (s *MigrationService) dropUnusedColumns(dst any) error {
dbCols, err := s.db.Migrator().ColumnTypes(dst)
if err != nil {
return err
// TableMigration 数据表迁移入口:先全量同步 schema,再增量删除字段并迁移数据
func (s *MigrationService) TableMigration() {
s.fullTableMigration()
s.fixTableConstraints()
s.incrementalTableMigration()
s.migrateChatAppSystemPromptFromJSON()
}
// migrateChatAppSystemPromptFromJSON 将智能体 system_prompt 字段中历史 JSON 数组
// 解析后取出 role 为 system 的 content,覆盖回 system_prompt(纯文本)
func (s *MigrationService) migrateChatAppSystemPromptFromJSON() {
key := "migrate:chat_app_system_prompt_json"
if s.redisClient.Get(context.Background(), key).Val() == "1" {
logger.Info("ChatApp system_prompt JSON 已迁移,跳过")
return
}
modelCols, err := s.modelColumnNames(dst)
if err != nil {
return err
logger.Info("开始迁移智能体 system_prompt 历史 JSON 数据...")
var apps []model.ChatApp
if err := s.db.Find(&apps).Error; err != nil {
logger.Errorf("查询 ChatApp 失败: %v", err)
return
}
for _, col := range dbCols {
name := col.Name()
if s.isProtectedColumn(name) {
updated := 0
for i := range apps {
raw := strings.TrimSpace(apps[i].SystemPrompt)
if raw == "" {
continue
}
if _, ok := modelCols[strings.ToLower(name)]; ok {
if len(raw) < 2 || raw[0] != '[' {
continue
}
logger.Infof("删除多余字段: %s.%s", s.tableName(dst), name)
if err := s.db.Migrator().DropColumn(dst, name); err != nil {
return fmt.Errorf("DropColumn %s: %w", name, err)
}
}
return nil
}
func (s *MigrationService) modelColumnNames(dst any) (map[string]struct{}, error) {
parsed, err := schema.Parse(dst, &schemaCache, s.db.Config.NamingStrategy)
if err != nil {
return nil, err
}
cols := make(map[string]struct{}, len(parsed.Fields))
for _, field := range parsed.Fields {
if field.DBName == "" || field.IgnoreMigration {
var messages []types.Message
if err := json.Unmarshal([]byte(raw), &messages); err != nil {
continue
}
cols[strings.ToLower(field.DBName)] = struct{}{}
var systemContent string
for _, m := range messages {
if strings.ToLower(strings.TrimSpace(m.Role)) == "system" && m.Content != "" {
systemContent = m.Content
break
}
}
if err := s.db.Model(&model.ChatApp{}).Where("id = ?", apps[i].Id).Update("system_prompt", systemContent).Error; err != nil {
logger.Warnf("更新 ChatApp id=%d system_prompt 失败: %v", apps[i].Id, err)
continue
}
updated++
}
return cols, nil
}
func (s *MigrationService) isProtectedColumn(name string) bool {
_, ok := protectedColumns[strings.ToLower(name)]
return ok
logger.Infof("智能体 system_prompt JSON 迁移完成,共更新 %d 条", updated)
s.redisClient.Set(context.Background(), key, "1", 0)
}
func (s *MigrationService) tableName(dst any) string {
stmt := &gorm.Statement{DB: s.db}
if err := stmt.Parse(dst); err != nil {
return fmt.Sprintf("%T", dst)
}
return stmt.Schema.Table
}
// schema.Parse 进程内复用的 schema cache
var schemaCache sync.Map
// 迁移配置数据
func (s *MigrationService) MigrateConfig(config *types.AppConfig) error {
@@ -374,6 +593,14 @@ func (s *MigrationService) migrateCommunicationConfig(config *types.AppConfig) e
"sign": config.SMS.Bao.Sign,
"code_template": config.SMS.Bao.CodeTemplate,
},
"tencent": map[string]any{
"secret_id": config.SMS.Tencent.SecretId,
"secret_key": config.SMS.Tencent.SecretKey,
"sms_sdk_app_id": config.SMS.Tencent.SmsSdkAppId,
"sign": config.SMS.Tencent.Sign,
"code_temp_id": config.SMS.Tencent.CodeTempId,
"region": config.SMS.Tencent.Region,
},
}
return s.saveConfig(types.ConfigKeySms, smsConfig)
}
@@ -406,3 +633,73 @@ func (s *MigrationService) saveConfig(key string, config any) error {
logger.Infof("成功迁移配置 %s", key)
return nil
}
// migrateSunoJobData 合并 suno_job 数据:从 Output 原始数据中解析出 tags 和 model_name 填入 params 字段
func (s *MigrationService) migrateSunoJobData() {
key := "migrate:suno_job_data"
if s.redisClient.Get(context.Background(), key).Val() == "1" {
logger.Info("SunoJob 数据已合并,跳过迁移")
return
}
logger.Info("开始合并 SunoJob 数据...")
// 查询所有有 output 数据的记录
var jobs []model.SunoJob
if err := s.db.Where("output != ? AND output != ''", "").Find(&jobs).Error; err != nil {
logger.Errorf("查询 SunoJob 数据失败: %v", err)
return
}
updatedCount := 0
for _, job := range jobs {
if job.Output == "" {
continue
}
// 解析 Output JSON 数据
var outputData struct {
Metadata struct {
Tags string `json:"tags"`
} `json:"metadata"`
ModelName string `json:"model_name"`
}
if err := json.Unmarshal([]byte(job.Output), &outputData); err != nil {
logger.Warnf("解析 Output 数据失败 (ID: %d): %v", job.Id, err)
continue
}
// 检查是否需要更新 params
needUpdate := false
params := job.Params
// 如果 params 中的 tags 为空,但 output 中有 tags,则更新
if params.Tags == "" && outputData.Metadata.Tags != "" {
params.Tags = outputData.Metadata.Tags
// 修复 tags 字段过长导致更新失败
if len(params.Tags) > 255 {
params.Tags = params.Tags[:255]
}
needUpdate = true
}
// 如果 params 中的 model 为空,但 output 中有 model_name,则更新
if params.Model == "" && outputData.ModelName != "" {
params.Model = outputData.ModelName
needUpdate = true
}
// 如果需要更新,则保存
if needUpdate {
if err := s.db.Model(&model.SunoJob{}).Where("id = ?", job.Id).Update("params", params).Error; err != nil {
logger.Errorf("更新 SunoJob 数据失败 (ID: %d): %v", job.Id, err)
continue
}
updatedCount++
}
}
logger.Infof("SunoJob 数据合并完成,共更新 %d 条记录", updatedCount)
s.redisClient.Set(context.Background(), key, "1", 0)
}
+28 -5
View File
@@ -12,17 +12,20 @@ import (
"errors"
"fmt"
"geekai/core/types"
logger2 "geekai/logger"
"geekai/log"
"geekai/store/model"
"geekai/utils"
"github.com/imroc/req/v3"
"gorm.io/gorm"
"io"
"time"
"github.com/imroc/req/v3"
"gorm.io/gorm"
"github.com/gin-gonic/gin"
)
var logger = log.GetLogger()
// Client MidJourney client
type Client struct {
client *req.Client
@@ -73,8 +76,6 @@ type QueryRes struct {
SubmitTime int `json:"submitTime"`
}
var logger = logger2.GetLogger()
func NewClient(db *gorm.DB) *Client {
return &Client{
client: req.C().SetTimeout(time.Minute).SetUserAgent("Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/123.0.0.0 Safari/537.36"),
@@ -182,6 +183,28 @@ func (c *Client) Variation(task types.MjTask) (ImageRes, error) {
return c.doRequest(body, apiPath, task.ChannelId)
}
// Modal 提交局部重绘(inpaint/ ZOOM,请求体 taskId 必填(提交成功返回的 taskId,不是查询结果里的 messageId),prompt、maskBase64 可选
func (c *Client) Modal(task types.MjTask) (ImageRes, error) {
apiPath := fmt.Sprintf("mj-%s/mj/submit/modal", task.Mode)
taskId := task.TaskId
if taskId == "" {
taskId = task.MessageId
}
if taskId == "" {
return ImageRes{}, fmt.Errorf("modal 任务缺少原图 taskId(提交成功返回的 ID")
}
body := map[string]any{
"taskId": taskId,
}
if task.Prompt != "" {
body["prompt"] = task.Prompt
}
if task.MaskBase64 != "" {
body["maskBase64"] = task.MaskBase64
}
return c.doRequest(body, apiPath, task.ChannelId)
}
func (c *Client) doRequest(body interface{}, apiPath string, channel string) (ImageRes, error) {
var res ImageRes
session := c.db.Session(&gorm.Session{}).Where("type", "mj").Where("enabled", true)
+3
View File
@@ -97,6 +97,9 @@ func (s *Service) Run() {
case types.TaskSwapFace:
res, err = s.client.SwapFace(task)
break
case types.TaskModal:
res, err = s.client.Modal(task)
break
}
if err != nil || (res.Code != 1 && res.Code != 22) {
@@ -2,8 +2,6 @@ package moderation
import (
"geekai/core/types"
logger2 "geekai/logger"
)
// * +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
@@ -13,8 +11,6 @@ import (
// * @Author yangjian102621@163.com
// * +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
var logger = logger2.GetLogger()
type Service interface {
Moderate(text string) (types.ModerationResult, error)
}
+167
View File
@@ -0,0 +1,167 @@
package oss
// * +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
// * 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 (
"bytes"
"context"
"encoding/base64"
"fmt"
"geekai/core/types"
"geekai/utils"
"net/http"
"net/url"
"path/filepath"
"strings"
"time"
"github.com/gin-gonic/gin"
"github.com/tencentyun/cos-go-sdk-v5"
)
type TencentOss struct {
config types.TencentOssConfig
client *cos.Client
proxyURL string
}
func NewTencentOss(sysConfig *types.SystemConfig, appConfig *types.AppConfig) (*TencentOss, error) {
s := &TencentOss{
proxyURL: appConfig.ProxyURL,
}
err := s.UpdateConfig(sysConfig.OSS.Tencent)
if err != nil {
logger.Warnf("腾讯云COS初始化失败: %v", err)
}
return s, nil
}
func (s *TencentOss) UpdateConfig(config types.TencentOssConfig) error {
if config.Bucket == "" || config.Region == "" || config.SecretId == "" || config.SecretKey == "" {
// 配置不完整时不初始化客户端
s.config = config
return nil
}
// 构建 COS 客户端 URL
cosURL := fmt.Sprintf("https://%s.cos.%s.myqcloud.com", config.Bucket, config.Region)
u, err := url.Parse(cosURL)
if err != nil {
return fmt.Errorf("error parsing COS URL: %v", err)
}
// 创建 COS 客户端
b := &cos.BaseURL{BucketURL: u}
client := cos.NewClient(b, &http.Client{
Transport: &cos.AuthorizationTransport{
SecretID: config.SecretId,
SecretKey: config.SecretKey,
},
})
s.client = client
s.config = config
return nil
}
func (s TencentOss) PutFile(ctx *gin.Context, name string) (File, error) {
// 解析表单
file, err := ctx.FormFile(name)
if err != nil {
return File{}, err
}
// 打开上传文件
src, err := file.Open()
if err != nil {
return File{}, err
}
defer src.Close()
fileExt := filepath.Ext(file.Filename)
objectKey := fmt.Sprintf("%d%s", time.Now().UnixMicro(), fileExt)
// 上传文件
_, err = s.client.Object.Put(ctx, objectKey, src, nil)
if err != nil {
return File{}, err
}
// 生成文件 URL
fileURL := s.generateURL(objectKey)
return File{
Name: file.Filename,
ObjKey: objectKey,
URL: fileURL,
Ext: fileExt,
Size: file.Size,
}, nil
}
func (s TencentOss) PutUrlFile(fileURL string, ext string, useProxy bool) (string, error) {
var fileData []byte
var err error
if useProxy {
fileData, err = utils.DownloadImage(fileURL, s.proxyURL)
} else {
fileData, err = utils.DownloadImage(fileURL, "")
}
if err != nil {
return "", fmt.Errorf("error with download image: %v", err)
}
parse, err := url.Parse(fileURL)
if err != nil {
return "", fmt.Errorf("error with parse image URL: %v", err)
}
if ext == "" {
ext = filepath.Ext(parse.Path)
}
objectKey := fmt.Sprintf("%d%s", time.Now().UnixMicro(), ext)
// 上传文件字节数据
_, err = s.client.Object.Put(context.Background(), objectKey, bytes.NewReader(fileData), nil)
if err != nil {
return "", err
}
return s.generateURL(objectKey), nil
}
func (s TencentOss) PutBase64(base64Img string) (string, error) {
imageData, err := base64.StdEncoding.DecodeString(base64Img)
if err != nil {
return "", fmt.Errorf("error decoding base64:%v", err)
}
objectKey := fmt.Sprintf("%d.png", time.Now().UnixMicro())
// 上传文件字节数据
_, err = s.client.Object.Put(context.Background(), objectKey, bytes.NewReader(imageData), nil)
if err != nil {
return "", err
}
return s.generateURL(objectKey), nil
}
func (s TencentOss) Delete(fileURL string) error {
var objectKey string
if strings.HasPrefix(fileURL, "http") {
objectKey = filepath.Base(fileURL)
} else {
objectKey = fileURL
}
_, err := s.client.Object.Delete(context.Background(), objectKey)
return err
}
// generateURL 生成文件访问 URL
func (s TencentOss) generateURL(objectKey string) string {
if s.config.Domain != "" {
// 使用自定义域名
return fmt.Sprintf("%s/%s", strings.TrimSuffix(s.config.Domain, "/"), objectKey)
}
// 使用 COS 默认域名
return fmt.Sprintf("https://%s.cos.%s.myqcloud.com/%s", s.config.Bucket, s.config.Region, objectKey)
}
var _ Uploader = TencentOss{}
+1
View File
@@ -13,6 +13,7 @@ const Local = "local"
const Minio = "minio"
const QiNiu = "qiniu"
const AliYun = "aliyun"
const Tencent = "tencent"
type File struct {
Name string `json:"name"`
+89 -13
View File
@@ -8,32 +8,41 @@ package oss
// * +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
import (
"fmt"
"geekai/core/types"
"strings"
logger2 "geekai/logger"
"geekai/log"
)
var logger = logger2.GetLogger()
var logger = log.GetLogger()
// 默认缩略图模板(本地存储格式)
const DefaultThumbTemplate = "?imageView2/4/w/{width}/h/{height}/q/75"
type UploaderManager struct {
local *LocalStorage
aliyun *AliYunOss
mini *MiniOss
qiniu *QiNiuOss
active string
local *LocalStorage
aliyun *AliYunOss
mini *MiniOss
qiniu *QiNiuOss
tencent *TencentOss
active string
ossConfig types.OSSConfig // 保存当前OSS配置
}
func NewUploaderManager(sysConfig *types.SystemConfig, local *LocalStorage, aliyun *AliYunOss, mini *MiniOss, qiniu *QiNiuOss) (*UploaderManager, error) {
func NewUploaderManager(sysConfig *types.SystemConfig, local *LocalStorage, aliyun *AliYunOss, mini *MiniOss, qiniu *QiNiuOss, tencent *TencentOss) (*UploaderManager, error) {
if sysConfig.OSS.Active == "" {
sysConfig.OSS.Active = Local
}
return &UploaderManager{
active: sysConfig.OSS.Active,
local: local,
aliyun: aliyun,
mini: mini,
qiniu: qiniu,
active: sysConfig.OSS.Active,
local: local,
aliyun: aliyun,
mini: mini,
qiniu: qiniu,
tencent: tencent,
ossConfig: sysConfig.OSS,
}, nil
}
@@ -47,6 +56,8 @@ func (m *UploaderManager) GetUploadHandler() Uploader {
return m.mini
case QiNiu:
return m.qiniu
case Tencent:
return m.tencent
}
return m.local
}
@@ -61,6 +72,71 @@ func (m *UploaderManager) UpdateConfig(config types.OSSConfig) {
m.mini.UpdateConfig(config.Minio)
case QiNiu:
m.qiniu.UpdateConfig(config.QiNiu)
case Tencent:
m.tencent.UpdateConfig(config.Tencent)
}
m.active = config.Active
m.ossConfig = config
}
// GetThumbURL 根据原始图片URL和尺寸生成缩略图URL
// 如果模板为空,使用默认模板(本地存储格式)
// 如果明确设置为空字符串(表示不支持缩略图),返回原图URL
func (m *UploaderManager) GetThumbURL(originalURL string, width, height int) string {
var template string
// 根据当前激活的存储引擎获取对应的模板
switch m.active {
case Local:
template = m.ossConfig.Local.ThumbTemplate
case AliYun:
template = m.ossConfig.AliYun.ThumbTemplate
case Minio:
template = m.ossConfig.Minio.ThumbTemplate
case QiNiu:
template = m.ossConfig.QiNiu.ThumbTemplate
case Tencent:
template = m.ossConfig.Tencent.ThumbTemplate
default:
template = m.ossConfig.Local.ThumbTemplate
}
// 如果模板为空,使用默认模板(兼容旧配置)
if template == "" {
template = DefaultThumbTemplate
}
// 替换变量
thumbURL := strings.ReplaceAll(template, "{width}", fmt.Sprintf("%d", width))
thumbURL = strings.ReplaceAll(thumbURL, "{height}", fmt.Sprintf("%d", height))
// 拼接原始URL和缩略图参数
return originalURL + thumbURL
}
// GetThumbTemplate 获取当前存储引擎的缩略图模板
func (m *UploaderManager) GetThumbTemplate() string {
var template string
switch m.active {
case Local:
template = m.ossConfig.Local.ThumbTemplate
case AliYun:
template = m.ossConfig.AliYun.ThumbTemplate
case Minio:
template = m.ossConfig.Minio.ThumbTemplate
case QiNiu:
template = m.ossConfig.QiNiu.ThumbTemplate
case Tencent:
template = m.ossConfig.Tencent.ThumbTemplate
default:
template = m.ossConfig.Local.ThumbTemplate
}
// 如果模板为空,返回默认模板
if template == "" {
return DefaultThumbTemplate
}
return template
}
+2 -2
View File
@@ -11,7 +11,7 @@ import (
"context"
"fmt"
"geekai/core/types"
logger2 "geekai/logger"
"geekai/log"
"net/http"
"os"
@@ -24,7 +24,7 @@ type AlipayService struct {
config *types.AlipayConfig
}
var logger = logger2.GetLogger()
var logger = log.GetLogger()
func NewAlipayService(sysConfig *types.SystemConfig) (*AlipayService, error) {
config := sysConfig.Payment.Alipay
+13 -1
View File
@@ -9,6 +9,7 @@ package payment
import (
"context"
"encoding/json"
"fmt"
"geekai/core/types"
"geekai/utils"
@@ -88,7 +89,18 @@ func (s *WxPayService) Pay(params PayRequest) (string, error) {
if wxRsp.Code != wechat.Success {
return "", fmt.Errorf("error status with generating pay url: %v", wxRsp.Error)
}
return wxRsp.Response.PrepayId, nil
// 签名
payParams, err := s.client.PaySignOfJSAPI(s.config.AppId, wxRsp.Response.PrepayId)
if err != nil {
return "", fmt.Errorf("error with generating jsapi pay sign: %v", err)
}
payParamsBytes, err := json.Marshal(payParams)
if err != nil {
return "", fmt.Errorf("error with marshaling pay params: %v", err)
}
return string(payParamsBytes), nil
} else if params.Device == "pc" {
wxRsp, err := s.client.V3TransactionNative(context.Background(), bm)
if err != nil {
Binary file not shown.
+404
View File
@@ -0,0 +1,404 @@
package ppt
import (
"bytes"
"context"
_ "embed"
"fmt"
"geekai/core/types"
"image"
_ "image/gif"
"io"
"net"
"net/http"
"net/url"
"os"
"path/filepath"
"sort"
"strings"
"time"
"unicode"
"github.com/jung-kurt/gofpdf/v2"
"github.com/ktye/pptx"
_ "golang.org/x/image/webp"
)
//go:embed embed/minimal.pptx
var minimalPptxTemplate []byte
const (
exportHTTPTimeout = 60 * time.Second
// 必须与 embed/minimal.pptx 中 p:sldSz 一致(当前模板为 4:310\"×7.5\"
slideEmuW pptx.Dimension = 9144000
slideEmuH pptx.Dimension = 6858000
)
// ExportFormat 导出类型
type ExportFormat string
const (
ExportFormatPDF ExportFormat = "pdf"
ExportFormatPPTX ExportFormat = "pptx"
)
// ExportMimeType 返回 Content-Type
func ExportMimeType(f ExportFormat) string {
switch f {
case ExportFormatPDF:
return "application/pdf"
case ExportFormatPPTX:
return "application/vnd.openxmlformats-officedocument.presentationml.presentation"
default:
return "application/octet-stream"
}
}
// ExportFileExt 返回文件扩展名(含点)
func ExportFileExt(f ExportFormat) string {
switch f {
case ExportFormatPDF:
return ".pdf"
case ExportFormatPPTX:
return ".pptx"
default:
return ""
}
}
// ParseExportFormat 解析 query format
func ParseExportFormat(s string) (ExportFormat, bool) {
switch strings.ToLower(strings.TrimSpace(s)) {
case "pdf":
return ExportFormatPDF, true
case "pptx", "ppt":
return ExportFormatPPTX, true
default:
return "", false
}
}
// SanitizeExportBaseName 用于下载文件名的主体(不含扩展名)
func SanitizeExportBaseName(title, taskID string) string {
s := strings.TrimSpace(title)
repl := strings.NewReplacer(
"/", "_", "\\", "_", ":", "_", "*", "_", "?", "_", "\"", "_", "<", "_", ">", "_", "|", "_",
)
s = repl.Replace(s)
var b strings.Builder
for _, r := range s {
if r == unicode.ReplacementChar || r < 32 {
continue
}
b.WriteRune(r)
}
s = strings.TrimSpace(b.String())
if s == "" {
s = strings.TrimSpace(taskID)
}
if len([]rune(s)) > 120 {
rs := []rune(s)
s = string(rs[:120])
}
return s
}
// ContentDispositionAttachment RFC 5987,兼容旧客户端
func ContentDispositionAttachment(filename string) string {
ascii := filename
for _, r := range filename {
if r > 127 || r == '"' || r == '\\' {
ascii = "export" + strings.ToLower(filepath.Ext(filename))
if ascii == "export" {
ascii = "export.bin"
}
break
}
}
return fmt.Sprintf(`attachment; filename="%s"; filename*=UTF-8''%s`, ascii, url.PathEscape(filename))
}
// mapLocalUploadFile 将站点相对路径或完整 BaseURL 前缀映射为本地文件路径(local OSS)
func mapLocalUploadFile(raw string, local types.LocalStorageConfig) (string, bool) {
raw = strings.TrimSpace(raw)
if raw == "" || local.BasePath == "" {
return "", false
}
bp := filepath.Clean(local.BasePath)
bu := strings.TrimSuffix(strings.TrimSpace(local.BaseURL), "/")
if bu != "" && strings.HasPrefix(raw, bu) {
suffix := strings.TrimPrefix(strings.TrimPrefix(raw, bu), "/")
return filepath.Join(bp, suffix), true
}
if strings.HasPrefix(raw, local.BaseURL) {
return filepath.Join(bp, strings.TrimPrefix(raw, local.BaseURL)), true
}
return "", false
}
func originFromBaseURL(baseURL string) string {
u, err := url.Parse(strings.TrimSpace(baseURL))
if err != nil || u.Scheme == "" || u.Host == "" {
return ""
}
return u.Scheme + "://" + u.Host
}
func defaultOriginFromListen(listen string) string {
listen = strings.TrimSpace(listen)
if listen == "" {
return ""
}
host, port, err := net.SplitHostPort(listen)
if err != nil {
if strings.HasPrefix(listen, ":") {
return "http://127.0.0.1" + listen
}
return ""
}
if host == "0.0.0.0" || host == "::" || host == "" {
host = "127.0.0.1"
}
return "http://" + net.JoinHostPort(host, port)
}
// resolveAbsoluteImageURL 将可能为相对路径的地址转为可 HTTP 访问的绝对 URL
func resolveAbsoluteImageURL(raw string, local types.LocalStorageConfig, app *types.AppConfig) string {
raw = strings.TrimSpace(raw)
if raw == "" {
return raw
}
if u, err := url.Parse(raw); err == nil && u.Scheme != "" && u.Host != "" {
return raw
}
if strings.HasPrefix(raw, "//") {
return "https:" + raw
}
origin := originFromBaseURL(local.BaseURL)
if origin == "" && app != nil {
origin = originFromBaseURL(app.StaticUrl)
}
if origin == "" && app != nil {
origin = defaultOriginFromListen(app.Listen)
}
if origin == "" {
return raw
}
if strings.HasPrefix(raw, "/") {
return strings.TrimSuffix(origin, "/") + raw
}
return strings.TrimSuffix(origin, "/") + "/" + raw
}
// BuildExportBytes 按幻灯片顺序拉取图片并生成 PDF 或 PPTX(oss/app 用于解析相对路径图片 URL)
func BuildExportBytes(ctx context.Context, slides []SlideData, format ExportFormat, oss types.OSSConfig, app *types.AppConfig) ([]byte, error) {
raws, err := fetchSlideImages(ctx, slides, oss, app)
if err != nil {
return nil, err
}
if len(raws) == 0 {
return nil, fmt.Errorf("没有可导出的幻灯片图片")
}
switch format {
case ExportFormatPDF:
return buildPDF(raws)
case ExportFormatPPTX:
return buildPPTX(raws)
default:
return nil, fmt.Errorf("不支持的导出格式")
}
}
func fetchSlideImages(ctx context.Context, slides []SlideData, oss types.OSSConfig, app *types.AppConfig) ([][]byte, error) {
cp := append([]SlideData(nil), slides...)
sort.Slice(cp, func(i, j int) bool { return cp[i].SlideIndex < cp[j].SlideIndex })
client := &http.Client{Timeout: exportHTTPTimeout}
var out [][]byte
for _, s := range cp {
u := strings.TrimSpace(s.ImageURL)
if u == "" {
continue
}
var body []byte
if oss.Active == "local" {
if fp, ok := mapLocalUploadFile(u, oss.Local); ok {
b, err := os.ReadFile(fp)
if err == nil {
if _, err := decodeImageBytes(b); err == nil {
body = b
}
}
}
}
if len(body) == 0 {
absURL := resolveAbsoluteImageURL(u, oss.Local, app)
req, err := http.NewRequestWithContext(ctx, http.MethodGet, absURL, nil)
if err != nil {
return nil, fmt.Errorf("幻灯片 %d: %w", s.SlideIndex, err)
}
resp, err := client.Do(req)
if err != nil {
return nil, fmt.Errorf("幻灯片 %d 下载失败: %w", s.SlideIndex, err)
}
body, err = io.ReadAll(io.LimitReader(resp.Body, 32<<20))
_ = resp.Body.Close()
if err != nil {
return nil, fmt.Errorf("幻灯片 %d 读取失败: %w", s.SlideIndex, err)
}
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
return nil, fmt.Errorf("幻灯片 %d 下载失败: HTTP %d", s.SlideIndex, resp.StatusCode)
}
}
if _, err := decodeImageBytes(body); err != nil {
return nil, fmt.Errorf("幻灯片 %d 不是有效图片: %v", s.SlideIndex, err)
}
out = append(out, body)
}
return out, nil
}
func decodeImageBytes(b []byte) (image.Image, error) {
m, _, err := image.Decode(bytes.NewReader(b))
return m, err
}
// fitSlideBoundsEMU 在幻灯片 EMU 框(与模板 p:sldSz 一致)内按原图比例 contain 居中(不裁切)
func fitSlideBoundsEMU(imgW, imgH int) (x, y, w, h pptx.Dimension) {
if imgW <= 0 || imgH <= 0 {
return 0, 0, slideEmuW, slideEmuH
}
sw := float64(slideEmuW)
sh := float64(slideEmuH)
iw := float64(imgW)
ih := float64(imgH)
scale := sw / iw
if ih*scale > sh {
scale = sh / ih
}
wf := iw * scale
hf := ih * scale
w = pptx.Dimension(wf + 0.5)
h = pptx.Dimension(hf + 0.5)
x = pptx.Dimension((sw-wf)*0.5 + 0.5)
y = pptx.Dimension((sh-hf)*0.5 + 0.5)
return x, y, w, h
}
func buildPDF(images [][]byte) ([]byte, error) {
pdf := gofpdf.New("L", "mm", "A4", "")
pdf.SetMargins(0, 0, 0)
pdf.SetAutoPageBreak(false, 0)
// 像素 → mm(按 96 DPI),再按页面对比缩放以 contain 放入整页
const pxPerMM = 96.0 / 25.4
for i, raw := range images {
im, err := decodeImageBytes(raw)
if err != nil {
return nil, fmt.Errorf("第 %d 页: %w", i+1, err)
}
b := im.Bounds()
pxW := float64(b.Dx())
pxH := float64(b.Dy())
if pxW <= 0 || pxH <= 0 {
return nil, fmt.Errorf("第 %d 页: 图片尺寸无效", i+1)
}
pdf.AddPage()
pageW, pageH := pdf.GetPageSize()
imgWmm := pxW / pxPerMM
imgHmm := pxH / pxPerMM
scale := pageW / imgWmm
if imgHmm*scale > pageH {
scale = pageH / imgHmm
}
w := imgWmm * scale
h := imgHmm * scale
x := (pageW - w) / 2
y := (pageH - h) / 2
name := fmt.Sprintf("slide%d", i)
opt := gofpdf.ImageOptions{ReadDpi: false}
tp := sniffImageType(raw)
if tp != "" {
opt.ImageType = tp
}
if pdf.RegisterImageOptionsReader(name, opt, bytes.NewReader(raw)) == nil {
return nil, fmt.Errorf("第 %d 页: 无法写入 PDF 图片", i+1)
}
pdf.ImageOptions(name, x, y, w, h, false, opt, 0, "")
}
var buf bytes.Buffer
if err := pdf.Output(&buf); err != nil {
return nil, err
}
return buf.Bytes(), nil
}
func sniffImageType(b []byte) string {
if len(b) < 12 {
return ""
}
switch {
case len(b) >= 2 && b[0] == 0xFF && b[1] == 0xD8:
return "jpg"
case len(b) >= 8 && string(b[0:8]) == "\x89PNG\r\n\x1a\n":
return "png"
case len(b) >= 6 && string(b[0:6]) == "GIF87a" || string(b[0:6]) == "GIF89a":
return "gif"
case len(b) >= 12 && string(b[0:4]) == "RIFF" && string(b[8:12]) == "WEBP":
return "webp"
default:
return ""
}
}
func buildPPTX(images [][]byte) ([]byte, error) {
if len(minimalPptxTemplate) == 0 {
return nil, fmt.Errorf("内置 PPT 模板缺失")
}
tmp, err := os.CreateTemp("", "ppt-export-*.pptx")
if err != nil {
return nil, err
}
path := tmp.Name()
if _, err := tmp.Write(minimalPptxTemplate); err != nil {
_ = tmp.Close()
_ = os.Remove(path)
return nil, err
}
if err := tmp.Close(); err != nil {
_ = os.Remove(path)
return nil, err
}
defer func() { _ = os.Remove(path) }()
f, err := pptx.Open(path)
if err != nil {
return nil, err
}
for _, raw := range images {
im, err := decodeImageBytes(raw)
if err != nil {
f.Abort()
return nil, err
}
b := im.Bounds()
ex, ey, ew, eh := fitSlideBoundsEMU(b.Dx(), b.Dy())
slide := pptx.Slide{
Images: []pptx.Image{
pptx.NewImage(im, ex, ey, ew, eh),
},
}
if err := f.Add(slide); err != nil {
f.Abort()
return nil, err
}
}
if err := f.Close(); err != nil {
return nil, err
}
return os.ReadFile(path)
}
+263
View File
@@ -0,0 +1,263 @@
package ppt
import (
"context"
"fmt"
"geekai/core/types"
"geekai/log"
"net/http"
"time"
"github.com/imroc/req/v3"
"github.com/volcengine/volcengine-go-sdk/service/arkruntime"
"github.com/volcengine/volcengine-go-sdk/service/arkruntime/model"
"github.com/volcengine/volcengine-go-sdk/volcengine"
"golang.org/x/time/rate"
)
var imageLogger = log.GetLogger()
// ImageGenerator 图片生成适配器接口
type ImageGenerator interface {
Provider() string
Generate(ctx context.Context, prompt string) (string, error)
// GenerateWithReference 图生图;referenceImages 为公网 URL 或 data:image/...;base64,...(本地文件应在调用前经 PrepareReferenceInputsForImg2Img 转换)
GenerateWithReference(ctx context.Context, prompt string, referenceImages []string) (string, error)
}
// Nano Banana 适配器(OpenAI DALL-E 风格 API
type nanoBananaImageGenerator struct {
client *req.Client
cfg types.PPTConfig
limiter *rate.Limiter
}
// Seedream 适配器(火山引擎 arkruntime SDK
type seedreamImageGenerator struct {
cfg types.PPTConfig
limiter *rate.Limiter
}
// NewImageGenerator 根据配置创建对应的图片生成适配器
func NewImageGenerator(cfg types.PPTConfig) (ImageGenerator, error) {
qps := cfg.QPSLimit
if qps <= 0 {
qps = 1
}
limiter := rate.NewLimiter(rate.Limit(qps), 1)
switch cfg.ActiveImageProvider {
case types.PPTImageProviderNanoBanana:
if cfg.NanoBananaApiURL == "" || cfg.NanoBananaApiKey == "" {
return nil, fmt.Errorf("nano banana api not configured")
}
return &nanoBananaImageGenerator{
client: req.C().SetTimeout(3 * time.Minute),
cfg: cfg,
limiter: limiter,
}, nil
case types.PPTImageProviderSeedream:
if cfg.SeedreamBaseURL == "" || cfg.SeedreamApiKey == "" || cfg.SeedreamModel == "" {
return nil, fmt.Errorf("seedream api not configured")
}
return &seedreamImageGenerator{
cfg: cfg,
limiter: limiter,
}, nil
default:
return nil, fmt.Errorf("unsupported image provider: %s", cfg.ActiveImageProvider)
}
}
func (g *nanoBananaImageGenerator) Provider() string {
return string(types.PPTImageProviderNanoBanana)
}
// nanoBananaReq 按 OpenAI DALL-E 风格 / Nano-banana API 文档
type nanoBananaReq struct {
Model string `json:"model"`
Prompt string `json:"prompt"`
ResponseFormat string `json:"response_format,omitempty"` // url 或 b64_json
AspectRatio string `json:"aspect_ratio,omitempty"` // 1:1, 4:3, 3:4, 16:9, 9:16, 2:3, 3:2, 4:5, 5:4, 21:9
Image []string `json:"image,omitempty"` // 参考图 url 或 b64
}
// nanoBananaRes 响应为 data[].urlDALL-E 风格)
type nanoBananaRes struct {
Data []struct {
URL string `json:"url,omitempty"`
B64JSON string `json:"b64_json,omitempty"`
} `json:"data"`
}
type nanoBananaErr struct {
Error struct {
Message string `json:"message"`
} `json:"error"`
}
func (g *nanoBananaImageGenerator) buildReqBody(prompt string, referenceImages []string) nanoBananaReq {
modelName := g.cfg.NanoBananaModel
if modelName == "" {
modelName = "nano-banana"
}
reqBody := nanoBananaReq{
Model: modelName,
Prompt: prompt,
}
if len(referenceImages) > 0 {
reqBody.Image = referenceImages
}
if g.cfg.NanoBananaResponseFormat != "" {
reqBody.ResponseFormat = g.cfg.NanoBananaResponseFormat
} else {
reqBody.ResponseFormat = "url"
}
if g.cfg.NanoBananaAspectRatio != "" {
reqBody.AspectRatio = g.cfg.NanoBananaAspectRatio
} else {
reqBody.AspectRatio = "16:9"
}
return reqBody
}
func (g *nanoBananaImageGenerator) Generate(ctx context.Context, prompt string) (string, error) {
return g.GenerateWithReference(ctx, prompt, nil)
}
func (g *nanoBananaImageGenerator) GenerateWithReference(ctx context.Context, prompt string, referenceImages []string) (string, error) {
reqBody := g.buildReqBody(prompt, referenceImages)
var (
result nanoBananaRes
errRes nanoBananaErr
)
do := func() (int, error) {
if err := g.limiter.Wait(ctx); err != nil {
return 0, err
}
imageLogger.Infof("nano banana generate image, api: %s", g.cfg.NanoBananaApiURL)
r, err := g.client.R().
SetContext(ctx).
SetHeader("Content-Type", "application/json").
SetHeader("Authorization", "Bearer "+g.cfg.NanoBananaApiKey).
SetBody(reqBody).
SetSuccessResult(&result).
SetErrorResult(&errRes).
Post(g.cfg.NanoBananaApiURL)
if err != nil {
return 0, err
}
if r.IsErrorState() {
return r.StatusCode, fmt.Errorf("nano banana error: %s, %s", r.Status, errRes.Error.Message)
}
if len(result.Data) == 0 || result.Data[0].URL == "" {
return r.StatusCode, fmt.Errorf("nano banana returned empty data")
}
return r.StatusCode, nil
}
if err := callWithRetry(ctx, do); err != nil {
return "", err
}
return result.Data[0].URL, nil
}
func (g *seedreamImageGenerator) Provider() string {
return string(types.PPTImageProviderSeedream)
}
func (g *seedreamImageGenerator) Generate(ctx context.Context, prompt string) (string, error) {
return g.GenerateWithReference(ctx, prompt, nil)
}
func (g *seedreamImageGenerator) GenerateWithReference(ctx context.Context, prompt string, referenceImages []string) (string, error) {
if err := g.limiter.Wait(ctx); err != nil {
return "", err
}
client := arkruntime.NewClientWithApiKey(g.cfg.SeedreamApiKey, arkruntime.WithBaseUrl(g.cfg.SeedreamBaseURL))
size := g.cfg.SeedreamSize
if size == "" {
size = "1920x1080"
}
responseFormat := g.cfg.SeedreamResponseType
if responseFormat == "" {
responseFormat = "url"
}
generateReq := model.GenerateImagesRequest{
Model: g.cfg.SeedreamModel,
Prompt: prompt,
Size: volcengine.String(size),
ResponseFormat: volcengine.String(responseFormat),
Watermark: volcengine.Bool(g.cfg.SeedreamWatermark),
}
if len(referenceImages) > 0 {
generateReq.Image = referenceImages
}
var lastErr error
for attempt := 0; attempt < 3; attempt++ {
if attempt > 0 {
select {
case <-ctx.Done():
return "", ctx.Err()
case <-time.After(time.Duration(attempt) * 2 * time.Second):
}
}
imageLogger.Infof("seedream generate image, api: %s", g.cfg.SeedreamBaseURL)
if err := generateReq.NormalizeImages(); err != nil {
return "", fmt.Errorf("seedream normalize images: %w", err)
}
resp, err := client.GenerateImages(ctx, generateReq)
if err != nil {
lastErr = fmt.Errorf("seedream error: %w", err)
continue
}
if resp.Data == nil || len(resp.Data) == 0 {
lastErr = fmt.Errorf("seedream returned empty data")
continue
}
if resp.Data[0].Url == nil || *resp.Data[0].Url == "" {
lastErr = fmt.Errorf("seedream returned empty url")
continue
}
return *resp.Data[0].Url, nil
}
return "", lastErr
}
// callWithRetry 对 429 错误做指数退避重试
func callWithRetry(ctx context.Context, fn func() (int, error)) error {
var (
retries = 3
backoffs = []time.Duration{2 * time.Second, 4 * time.Second, 8 * time.Second}
lastError error
)
for i := 0; i < retries; i++ {
status, err := fn()
if err == nil {
return nil
}
lastError = err
// 仅对 429 做指数退避重试
if status != http.StatusTooManyRequests || i == retries-1 {
break
}
select {
case <-ctx.Done():
return ctx.Err()
case <-time.After(backoffs[i]):
}
}
return lastError
}
+373
View File
@@ -0,0 +1,373 @@
package ppt
// * +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
// * 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"
"encoding/json"
"fmt"
"geekai/core/types"
"geekai/log"
"strings"
"time"
"github.com/imroc/req/v3"
)
var logger = log.GetLogger()
// slidePlan LLM 输出的分镜结构
type slidePlan struct {
SlideIndex int `json:"slide_index"`
Theme string `json:"theme"`
Title string `json:"title"`
Points []string `json:"points"`
ImagePrompt string `json:"image_prompt"`
}
// systemPrompt 图文并茂幻灯片:每页有插图且画面上含与主题一致的文字,文字量与生成模式相关
const systemPrompt = `
# Role
你是一位顶级的专业演示文稿(PPT)策划专家和 AI 图像提示词(Prompt)工程师。任务是根据用户提供的「内容大纲」或「设计要求」,生成一套逻辑清晰、视觉风格高度统一的幻灯片分镜数据。目标是生成**图文并茂**的幻灯片:每页既有文字又有插图,图片上直接呈现与本页内容一致的文字,而不是留白让用户后加文字。
# Rules
1. 全局风格锚定:根据大纲推断或遵循用户要求的全局视觉风格。所有配图必须严格遵循此风格。
2. 结构化拆解:合理拆分为多张幻灯片,单页最多 3-4 个简短要点。
3. 视觉转译(图文并茂):
- 为每页构思的 image_prompt 既要描述**插图画面**,也要明确**画面上应出现的文字**(如本页标题、要点或短句),与当页 theme、title、points 内容一致。不要描述留白或“用于排版文字的空间”。
- 图片中出现的所有文字必须使用与 theme、title、points **相同的语言**(即本次请求指定的输出语言)。
- 图片上文字的量由「生成模式」决定(见用户输入中的模式说明):
- **详细演示文稿**:图片上的说明文字可适当多一些,如本页要点、一两句说明。
- **演示用幻灯片**:图片上的文字尽量精简,如仅主标题或少量关键词,便于演讲时配合口述。
- image_prompt 须包含前缀「[全局风格描述]」,并清晰描述画面中的插图与文字内容(含具体要出现的文字及其语言),不要包含“不要在图片中生成任何文字”的约束。
4. 严格输出合法的纯 JSON 数组:
[
{"slide_index": 1, "theme": "...", "title": "...", "points": ["..."], "image_prompt": "..."}
]
禁止输出任何 Markdown 标记或多余文本。`
// notebookSystemPrompt 文档提炼:NotebookLM 风格输出 PPT 可用大纲文本(纯文本/Markdown)。
const notebookSystemPrompt = `
# Role
你是一位“NotebookLM 风格”的专业文档理解与提炼助手。任务是基于用户提供的「原始文档文本」和「设计要求」,提炼出可用于制作 PPT 的结构化大纲内容。
# Output Requirements
1. 输出必须是纯文本/Markdown(允许使用标题与列表),禁止输出任何 JSON。
2. 禁止输出代码块(不要出现代码块语法)。
3. 不要输出解释过程、不要复述提示词。
4. 大纲必须是“内容大纲/要点”,用于后续继续拆分成幻灯片,而不是直接输出最终幻灯片分镜。
# Rules
1. 文档优先:尽可能从原始文档中提取信息与措辞;若文档缺失关键点,则给出合理补全的“建议方向”,并明确标注为“(建议)”。
2. 贴合设计要求:根据设计要求调整大纲的语气、侧重点、术语风格,使内容更符合目标受众与整体风格。
3. 结构清晰:使用分层标题(例如:# 总主题、## 模块/章节、### 要点),并为每个模块给出 2-4 个要点句(可直接用于 PPT 每页标题/要点)。
4. 语言一致:所有输出语言必须与本次请求指定的语言一致。
`
// LLMClient 分镜 LLM 客户端
type LLMClient struct {
httpClient *req.Client
}
func NewLLMClient() *LLMClient {
return &LLMClient{
httpClient: req.C().SetTimeout(2 * time.Minute),
}
}
// GenerateSlides 调用大模型生成分镜列表。language 约束输出语言,mode 约束图中文字量,maxPages 约束恰好生成 N 页。
func (c *LLMClient) GenerateSlides(ctx context.Context, cfg types.PPTConfig, content, prompt, language, mode string, maxPages int) ([]slidePlan, error) {
if cfg.OutlineLLMApiURL == "" {
return nil, fmt.Errorf("outline LLM api url is empty")
}
if cfg.OutlineLLMApiKey == "" {
return nil, fmt.Errorf("outline LLM api key is empty")
}
if maxPages <= 0 {
maxPages = 10
}
if mode != "detailed" && mode != "slides" {
mode = "slides"
}
type message struct {
Role string `json:"role"`
Content string `json:"content"`
}
// 动态 system:加入页数约束
systemContent := systemPrompt + fmt.Sprintf("\n\n# 页数约束\n请将内容拆分为恰好 %d 页的幻灯片分镜,保证逻辑完整、故事线连贯,不要多也不要少。输出 JSON 数组长度必须为 %d。", maxPages, maxPages)
// 组装用户输入
userContent := fmt.Sprintf("下面是用户提供的演示文稿大纲内容:\n\n%s", content)
if prompt != "" {
userContent = fmt.Sprintf("%s\n\n额外的设计要求:%s", userContent, prompt)
}
if language != "" {
langHint := "中文"
if language == "en" || language == "en-US" {
langHint = "英文"
} else if language == "zh-CN" || language == "zh" {
langHint = "中文"
} else {
langHint = "语言代码 " + language + " 对应的语言"
}
userContent = fmt.Sprintf("%s\n\n请用%s输出所有分镜内容(theme、title、points、image_prompt 等均使用该语言;图片中出现的文字也必须是%s)。", userContent, langHint, langHint)
}
modeHint := "演示用幻灯片"
if mode == "detailed" {
modeHint = "详细演示文稿"
}
userContent = fmt.Sprintf("%s\n\n本次生成模式为:%s。请按上述规则控制每页插图中文字的量。", userContent, modeHint)
modelName := cfg.OutlineLLMModel
if modelName == "" {
modelName = "gpt-5.2"
}
body := map[string]any{
"model": modelName,
"messages": []message{
{Role: "user", Content: systemContent + "\n\n" + userContent},
},
"temperature": 0.8,
}
var respBody struct {
Choices []struct {
Message struct {
Content string `json:"content"`
} `json:"message"`
} `json:"choices"`
}
logger.Infof("generate PPT slides with outline LLM, api: %s", cfg.OutlineLLMApiURL)
r, err := c.httpClient.R().
SetContext(ctx).
SetHeader("Content-Type", "application/json").
SetHeader("Authorization", "Bearer "+cfg.OutlineLLMApiKey).
SetBody(body).
SetSuccessResult(&respBody).
Post(cfg.OutlineLLMApiURL)
if err != nil {
return nil, fmt.Errorf("request outline LLM failed: %v", err)
}
if r.IsErrorState() {
return nil, fmt.Errorf("outline LLM returned error status: %s", r.Status)
}
if len(respBody.Choices) == 0 {
return nil, fmt.Errorf("outline LLM returned empty choices")
}
contentStr := respBody.Choices[0].Message.Content
var plans []slidePlan
if err := json.Unmarshal([]byte(contentStr), &plans); err != nil {
return nil, fmt.Errorf("parse outline LLM json failed: %v, raw: %s", err, contentStr)
}
return plans, nil
}
// GenerateNotebookContent 调用文档提炼 LLM,把 rawDocText -> PPT 可用的 content(大纲/结构化要点)。
func (c *LLMClient) GenerateNotebookContent(ctx context.Context, cfg types.PPTConfig, rawDocText, designPrompt, language string) (string, error) {
if cfg.OutlineLLMApiURL == "" {
return "", fmt.Errorf("outline LLM api url is empty")
}
if cfg.OutlineLLMApiKey == "" {
return "", fmt.Errorf("outline LLM api key is empty")
}
rawDocText = strings.TrimSpace(rawDocText)
if rawDocText == "" {
return "", fmt.Errorf("rawDocText is empty")
}
// 对超长输入做保守截断,避免请求体过大或上下文溢出。
// 这里按“字符数”截断,真实 token 仍可能超出,但作为兜底足够。
const maxChars = 25000
runes := []rune(rawDocText)
if len(runes) > maxChars {
rawDocText = string(runes[:maxChars])
}
langHint := "中文"
if language == "en" || language == "en-US" {
langHint = "英文"
} else if language == "zh-CN" || language == "zh" {
langHint = "中文"
} else if language != "" {
langHint = "语言代码 " + language + " 对应的语言"
}
maxSlides := cfg.MaxSlidesPerTask
if maxSlides <= 0 {
maxSlides = 10
}
userContent := fmt.Sprintf("原始文档文本如下(可能很长):\n\n%s", rawDocText)
if strings.TrimSpace(designPrompt) != "" {
userContent = fmt.Sprintf("%s\n\n设计要求(风格/受众/侧重点等):\n%s", userContent, designPrompt)
}
userContent = fmt.Sprintf(
"%s\n\n请用%s输出 PPT 大纲内容。该大纲应便于拆分为不超过 %d 页的 PPT。",
userContent,
langHint,
maxSlides,
)
type message struct {
Role string `json:"role"`
Content string `json:"content"`
}
systemContent := notebookSystemPrompt + fmt.Sprintf("\n\n# 语言约束\n输出语言:%s。", langHint)
modelName := cfg.OutlineLLMModel
if modelName == "" {
modelName = "gpt-4o-mini"
}
body := map[string]any{
"model": modelName,
"messages": []message{
{Role: "user", Content: systemContent + "\n\n" + userContent},
},
"temperature": 0.4,
}
var respBody struct {
Choices []struct {
Message struct {
Content string `json:"content"`
} `json:"message"`
} `json:"choices"`
}
logger.Infof("generate PPT content with outline LLM, api: %s", cfg.OutlineLLMApiURL)
r, err := c.httpClient.R().
SetContext(ctx).
SetHeader("Content-Type", "application/json").
SetHeader("Authorization", "Bearer "+cfg.OutlineLLMApiKey).
SetBody(body).
SetSuccessResult(&respBody).
Post(cfg.OutlineLLMApiURL)
if err != nil {
return "", fmt.Errorf("request outline LLM failed: %v", err)
}
if r.IsErrorState() {
return "", fmt.Errorf("outline LLM returned error status: %s", r.Status)
}
if len(respBody.Choices) == 0 {
return "", fmt.Errorf("outline LLM returned empty choices")
}
return strings.TrimSpace(respBody.Choices[0].Message.Content), nil
}
// titleSystemPrompt 根据大纲生成用于任务列表的短标题(单行纯文本)
const titleSystemPrompt = `
# Role
你是「演示文稿命名助手」。用户会提供一份 PPT 内容大纲(可能含 Markdown)。请根据大纲主题与受众,生成一个**适合出现在任务列表中的短标题**。
# Output Rules
1. 只输出**一行**纯文本,不要换行、不要编号、不要引号包裹。
2. 长度建议 **20 个字以内**(中文)或 **8 个英文单词以内**;若大纲极长,仍只给概括性标题。
3. 输出语言必须与本次指定的「输出语言」一致。
4. 不要输出“标题:”“Title:”等前缀,不要复述本说明。
`
// GeneratePPTTitle 调用大模型根据 content 生成列表用短标题;失败时由调用方降级。
func (c *LLMClient) GeneratePPTTitle(ctx context.Context, cfg types.PPTConfig, content, language string) (string, error) {
if cfg.OutlineLLMApiURL == "" {
return "", fmt.Errorf("outline LLM api url is empty")
}
if cfg.OutlineLLMApiKey == "" {
return "", fmt.Errorf("outline LLM api key is empty")
}
content = strings.TrimSpace(content)
if content == "" {
return "", fmt.Errorf("content is empty")
}
const maxChars = 10000
runes := []rune(content)
if len(runes) > maxChars {
content = string(runes[:maxChars])
}
langHint := "中文"
if language == "en" || language == "en-US" {
langHint = "英文"
} else if language == "zh-CN" || language == "zh" {
langHint = "中文"
} else if language != "" {
langHint = "语言代码 " + language + " 对应的语言"
}
type message struct {
Role string `json:"role"`
Content string `json:"content"`
}
userContent := fmt.Sprintf("输出语言:%s。\n\n下面是用户提供的 PPT 大纲内容,请只返回列表标题:\n\n%s", langHint, content)
systemContent := titleSystemPrompt + fmt.Sprintf("\n\n# 语言约束\n请用%s撰写标题。", langHint)
modelName := cfg.OutlineLLMModel
if modelName == "" {
modelName = "gpt-4o-mini"
}
body := map[string]any{
"model": modelName,
"messages": []message{
{Role: "user", Content: systemContent + "\n\n" + userContent},
},
"temperature": 0.5,
}
var respBody struct {
Choices []struct {
Message struct {
Content string `json:"content"`
} `json:"message"`
} `json:"choices"`
}
logger.Infof("generate PPT list title with outline LLM, api: %s", cfg.OutlineLLMApiURL)
r, err := c.httpClient.R().
SetContext(ctx).
SetHeader("Content-Type", "application/json").
SetHeader("Authorization", "Bearer "+cfg.OutlineLLMApiKey).
SetBody(body).
SetSuccessResult(&respBody).
Post(cfg.OutlineLLMApiURL)
if err != nil {
return "", fmt.Errorf("request outline LLM failed: %v", err)
}
if r.IsErrorState() {
return "", fmt.Errorf("outline LLM returned error status: %s", r.Status)
}
if len(respBody.Choices) == 0 {
return "", fmt.Errorf("outline LLM returned empty choices")
}
raw := strings.TrimSpace(respBody.Choices[0].Message.Content)
if idx := strings.IndexAny(raw, "\r\n"); idx >= 0 {
raw = strings.TrimSpace(raw[:idx])
}
raw = strings.Trim(raw, `"'「」`)
return raw, nil
}
+884
View File
@@ -0,0 +1,884 @@
package ppt
// * +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
// * 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/types"
"geekai/service"
"geekai/service/oss"
"geekai/store/model"
"geekai/store/vo"
"geekai/utils"
"sort"
"strings"
"sync"
"time"
"golang.org/x/sync/errgroup"
"gorm.io/gorm"
)
// ErrInsufficientPower 用户算力不足以完成本次 PPT 任务(由 handler 映射文案)
var ErrInsufficientPower = errors.New("insufficient power for ppt task")
var (
// ErrPptTaskNotFound 表示任务不存在
ErrPptTaskNotFound = errors.New("ppt task not found")
// ErrPptTaskNotDeletable 表示任务状态不允许删除
ErrPptTaskNotDeletable = errors.New("ppt task not deletable")
// ErrPptTaskBusy 任务正在处理中
ErrPptTaskBusy = errors.New("ppt task is processing")
// ErrPptTaskNotResumable 无法继续生成(无分镜占位或已完成)
ErrPptTaskNotResumable = errors.New("ppt task cannot be resumed")
// ErrPptSlideNotFound 指定 slide_index 不存在
ErrPptSlideNotFound = errors.New("ppt slide not found")
// ErrPptSlideNoImage 该页尚无配图
ErrPptSlideNoImage = errors.New("ppt slide has no image")
// ErrPptInvalidVersionIndex 历史版本下标无效
ErrPptInvalidVersionIndex = errors.New("invalid slide version index")
)
// TaskStatus 任务状态
type TaskStatus string
const (
TaskStatusPending TaskStatus = "pending"
TaskStatusProcessing TaskStatus = "processing"
TaskStatusCompleted TaskStatus = "completed"
TaskStatusFailed TaskStatus = "failed"
)
// SlideData 单页 PPT 数据
type SlideData struct {
SlideIndex int `json:"slide_index"`
Theme string `json:"theme"`
Title string `json:"title"`
Points []string `json:"points"`
ImagePrompt string `json:"image_prompt"`
ImageURL string `json:"image_url"`
ImageHistory []vo.PPTSlideImageVersion `json:"image_history,omitempty"`
}
// Task PPT 生成任务(用于业务层与 API 返回,持久化在 DB)
type Task struct {
TaskID string `json:"task_id"`
UserID uint `json:"user_id"`
Status TaskStatus `json:"status"`
Content string `json:"content"`
Prompt string `json:"prompt"`
Language string `json:"language"`
Mode string `json:"mode"`
Pages int `json:"pages"`
Total int `json:"total_slides"`
Completed int `json:"completed_slides"`
Slides []SlideData `json:"slides"`
Title string `json:"title"`
Thumb string `json:"thumb"`
ErrorMessage string `json:"error_message"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
// TaskSummaryMap 返回任务摘要 map(公共字段),供 handler 补充独有字段后返回。
func (t *Task) TaskSummaryMap() map[string]any {
return map[string]any{
"task_id": t.TaskID,
"status": t.Status,
"total_slides": t.Total,
"completed_slides": t.Completed,
"created_at": t.CreatedAt.Unix(),
"updated_at": t.UpdatedAt.Unix(),
"title": t.Title,
"thumb": t.Thumb,
}
}
// Progress 任务进度信息
type Progress struct {
Total int `json:"total_slides"`
Completed int `json:"completed_slides"`
}
// PptService PPT 任务与生成流程(持久化、LLM 分镜、生图、转存、算力)
type PptService struct {
db *gorm.DB
userService *service.UserService
uploadManager *oss.UploaderManager
llm *LLMClient
// slidesLock 全局互斥:串行化 slides JSON 的写库,避免并发覆盖(写库极短,可接受排队)
slidesLock sync.Mutex
}
// NewPptService 创建 PptService
func NewPptService(db *gorm.DB, userService *service.UserService, uploadManager *oss.UploaderManager) *PptService {
return &PptService{
db: db,
userService: userService,
uploadManager: uploadManager,
llm: NewLLMClient(),
}
}
// GenerateNotebookContent 将原始文档文本提炼成 PPT 可用的大纲 content。
func (s *PptService) GenerateNotebookContent(ctx context.Context, cfg types.PPTConfig, rawDocText, designPrompt, language string) (string, error) {
if s.llm == nil {
s.llm = NewLLMClient()
}
return s.llm.GenerateNotebookContent(ctx, cfg, rawDocText, designPrompt, language)
}
func slideToVO(s SlideData) vo.PPTSlideData {
return vo.PPTSlideData{
SlideIndex: s.SlideIndex,
Theme: s.Theme,
Title: s.Title,
Points: s.Points,
ImagePrompt: s.ImagePrompt,
ImageURL: s.ImageURL,
ImageHistory: vo.PPTSlideImageVersions(s.ImageHistory),
}
}
func voToSlide(s vo.PPTSlideData) SlideData {
return SlideData{
SlideIndex: s.SlideIndex,
Theme: s.Theme,
Title: s.Title,
Points: s.Points,
ImagePrompt: s.ImagePrompt,
ImageURL: s.ImageURL,
ImageHistory: []vo.PPTSlideImageVersion(s.ImageHistory),
}
}
func voSlidesToBiz(slides vo.PPTSlides) []SlideData {
out := make([]SlideData, len(slides))
for i, sv := range slides {
out[i] = voToSlide(sv)
}
return out
}
func voSlidesToBizNormalized(slides vo.PPTSlides) []SlideData {
out := make([]SlideData, len(slides))
for i, sv := range slides {
sd := voToSlide(sv)
normalizeSlideImageHistory(&sd)
out[i] = sd
}
return out
}
// normalizeSlideImageHistory 旧数据仅有 image_url 时补全 image_history,便于前端展示历史
func normalizeSlideImageHistory(s *SlideData) {
if strings.TrimSpace(s.ImageURL) != "" && len(s.ImageHistory) == 0 {
s.ImageHistory = []vo.PPTSlideImageVersion{
{ImageURL: s.ImageURL, Prompt: strings.TrimSpace(s.ImagePrompt)},
}
}
}
// DerivePPTThumbFromSlides 按 slide_index 升序取第一张有图 URL(与 vo.PPTSlides 规则一致)
func DerivePPTThumbFromSlides(slides []SlideData) string {
if len(slides) == 0 {
return ""
}
cp := make([]SlideData, len(slides))
copy(cp, slides)
sort.Slice(cp, func(i, j int) bool {
return cp[i].SlideIndex < cp[j].SlideIndex
})
for _, s := range cp {
if strings.TrimSpace(s.ImageURL) != "" {
return s.ImageURL
}
}
return ""
}
func truncateTitleRunes(s string, max int) string {
if max <= 0 {
return s
}
r := []rune(s)
if len(r) <= max {
return s
}
return string(r[:max])
}
func taskToModel(t *Task) *model.PPTJob {
now := time.Now()
job := &model.PPTJob{
TaskId: t.TaskID,
UserId: t.UserID,
Status: string(t.Status),
ErrMsg: t.ErrorMessage,
Prompt: t.Prompt,
Title: t.Title,
Thumb: t.Thumb,
Content: t.Content,
Params: vo.PPTParams{Language: t.Language, Mode: t.Mode, Pages: t.Pages},
Slides: nil,
TotalSlides: t.Total,
CompletedSlides: t.Completed,
CreatedAt: now,
UpdatedAt: now,
}
if t.CreatedAt.IsZero() {
job.CreatedAt = now
job.UpdatedAt = now
} else {
job.CreatedAt = t.CreatedAt
job.UpdatedAt = t.UpdatedAt
}
return job
}
func modelToTask(j *model.PPTJob) *Task {
slides := voSlidesToBizNormalized(j.Slides)
return &Task{
TaskID: j.TaskId,
UserID: j.UserId,
Status: TaskStatus(j.Status),
Content: j.Content,
Prompt: j.Prompt,
Title: j.Title,
Thumb: j.Thumb,
Language: j.Params.Language,
Mode: j.Params.Mode,
Pages: j.Params.Pages,
Total: j.TotalSlides,
Completed: j.CompletedSlides,
Slides: slides,
ErrorMessage: j.ErrMsg,
CreatedAt: j.CreatedAt,
UpdatedAt: j.UpdatedAt,
}
}
// BuildPendingTask 校验算力与页数,组装待写入的 Task(未落库)
func (s *PptService) BuildPendingTask(taskID string, userID uint, userPower int, content, prompt, language, mode string, reqPages int) (*Task, types.PPTConfig, error) {
cfg, err := s.loadPPTConfig()
if err != nil {
return nil, cfg, err
}
// 目标页数:用户指定时取 min(请求页数, 服务端上限);未指定(0)时按服务端上限作为默认生成规模
effectivePages := cfg.MaxSlidesPerTask
if reqPages > 0 {
effectivePages = reqPages
if effectivePages > cfg.MaxSlidesPerTask {
effectivePages = cfg.MaxSlidesPerTask
}
}
estimatePower := effectivePages * cfg.PowerCostPerSlide
if estimatePower > 0 && userPower < estimatePower {
return nil, cfg, ErrInsufficientPower
}
effectiveMode := mode
if effectiveMode != "detailed" && effectiveMode != "slides" {
effectiveMode = "slides"
}
task := &Task{
TaskID: taskID,
UserID: userID,
Status: TaskStatusPending,
Content: content,
Prompt: prompt,
Language: language,
Pages: effectivePages,
Mode: effectiveMode,
}
return task, cfg, nil
}
// CreateTask 创建新任务并写入数据库(调用大模型生成列表标题后落库)
func (s *PptService) CreateTask(ctx context.Context, task *Task, cfg types.PPTConfig) error {
if s.llm == nil {
s.llm = NewLLMClient()
}
title, err := s.llm.GeneratePPTTitle(ctx, cfg, task.Content, task.Language)
if err != nil {
logger.Warnf("GeneratePPTTitle failed task_id=%s: %v", task.TaskID, err)
title = "未命名演示文稿"
} else {
title = strings.TrimSpace(title)
if title == "" {
title = "未命名演示文稿"
}
}
task.Title = truncateTitleRunes(title, 255)
task.CreatedAt = time.Now()
task.UpdatedAt = task.CreatedAt
task.Status = TaskStatusPending
job := taskToModel(task)
return s.db.Create(job).Error
}
// GetTask 从数据库获取任务
func (s *PptService) GetTask(taskID string) (*Task, bool) {
var job model.PPTJob
err := s.db.Where("task_id = ?", taskID).First(&job).Error
if err != nil || job.TaskId == "" {
return nil, false
}
return modelToTask(&job), true
}
// UpdateStatus 更新任务状态
func (s *PptService) UpdateStatus(taskID string, status TaskStatus) {
s.db.Model(&model.PPTJob{}).Where("task_id = ?", taskID).
Updates(map[string]interface{}{"status": string(status), "updated_at": time.Now()})
}
func countSlidesWithImage(slides []SlideData) int {
n := 0
for _, sl := range slides {
if strings.TrimSpace(sl.ImageURL) != "" {
n++
}
}
return n
}
func slidePlansToOutlines(plans []slidePlan) []SlideData {
out := make([]SlideData, len(plans))
for i, p := range plans {
out[i] = SlideData{
SlideIndex: p.SlideIndex,
Theme: p.Theme,
Title: p.Title,
Points: p.Points,
ImagePrompt: p.ImagePrompt,
ImageURL: "",
}
}
sort.Slice(out, func(i, j int) bool {
return out[i].SlideIndex < out[j].SlideIndex
})
return out
}
// saveSlidesOutline 分镜一出即落库:每页含 theme/title/points/image_promptimage_url 为空
func (s *PptService) saveSlidesOutline(taskID string, total int, slides []SlideData) error {
s.slidesLock.Lock()
defer s.slidesLock.Unlock()
voSlides := make(vo.PPTSlides, len(slides))
for i := range slides {
voSlides[i] = slideToVO(slides[i])
}
completed := countSlidesWithImage(slides)
return s.db.Model(&model.PPTJob{}).Where("task_id = ?", taskID).Updates(map[string]interface{}{
"slides": voSlides,
"total_slides": total,
"completed_slides": completed,
"updated_at": time.Now(),
}).Error
}
// ApplySlideImage 按 slide_index 原地写入 image_url,并刷新 completed_slides、thumb
func (s *PptService) ApplySlideImage(taskID string, slide SlideData) error {
s.slidesLock.Lock()
defer s.slidesLock.Unlock()
var job model.PPTJob
if err := s.db.Where("task_id = ?", taskID).First(&job).Error; err != nil {
return err
}
slides := job.Slides
found := false
for i := range slides {
if slides[i].SlideIndex == slide.SlideIndex {
slides[i].ImageURL = slide.ImageURL
if strings.TrimSpace(slide.ImageURL) != "" && len(slides[i].ImageHistory) == 0 {
slides[i].ImageHistory = vo.PPTSlideImageVersions{
{ImageURL: slide.ImageURL, Prompt: strings.TrimSpace(slide.ImagePrompt)},
}
}
found = true
break
}
}
if !found {
return fmt.Errorf("slide index %d not found", slide.SlideIndex)
}
job.Slides = slides
biz := voSlidesToBiz(slides)
return s.refreshJobMeta(&job, biz)
}
// refreshJobMeta 刷新 job 的 CompletedSlides/Thumb/UpdatedAt 并 Save。
// 调用前必须已持有 slidesLock。
func (s *PptService) refreshJobMeta(job *model.PPTJob, biz []SlideData) error {
job.CompletedSlides = countSlidesWithImage(biz)
job.Thumb = DerivePPTThumbFromSlides(biz)
job.UpdatedAt = time.Now()
return s.db.Save(job).Error
}
func (s *PptService) validateSlideOutline(task *Task) error {
if task.Total <= 0 {
return ErrPptTaskNotResumable
}
if len(task.Slides) < task.Total {
return ErrPptTaskNotResumable
}
seen := make(map[int]bool, len(task.Slides))
for _, sl := range task.Slides {
seen[sl.SlideIndex] = true
}
for i := 1; i <= task.Total; i++ {
if !seen[i] {
return ErrPptTaskNotResumable
}
}
return nil
}
func slidesNeedingImages(task *Task) []SlideData {
var need []SlideData
for _, sl := range task.Slides {
if strings.TrimSpace(sl.ImageURL) == "" {
need = append(need, sl)
}
}
sort.Slice(need, func(i, j int) bool {
return need[i].SlideIndex < need[j].SlideIndex
})
return need
}
func (s *PptService) userPower(userID uint) (int, error) {
var u model.User
if err := s.db.Where("id = ?", userID).First(&u).Error; err != nil {
return 0, err
}
return u.Power, nil
}
// runSlideImageJobs 为给定幻灯片列表并发生图(每张成功后 ApplySlideImage
func (s *PptService) runSlideImageJobs(ctx context.Context, task *Task, cfg types.PPTConfig, generator ImageGenerator, slides []SlideData) error {
if len(slides) == 0 {
return nil
}
if cfg.MaxConcurrentRequests <= 0 {
cfg.MaxConcurrentRequests = 3
}
group, ctx := errgroup.WithContext(ctx)
group.SetLimit(cfg.MaxConcurrentRequests)
for _, item := range slides {
slide := item
group.Go(func() error {
imgURL, err := generator.Generate(ctx, slide.ImagePrompt)
if err != nil {
return err
}
storedURL, err := s.uploadManager.GetUploadHandler().PutUrlFile(imgURL, ".png", false)
if err != nil {
return fmt.Errorf("转存图片失败:%w", err)
}
full := slide
full.ImageURL = storedURL
if err := s.ApplySlideImage(task.TaskID, full); err != nil {
return err
}
if cfg.PowerCostPerSlide > 0 {
err = s.userService.DecreasePower(task.UserID, cfg.PowerCostPerSlide, model.PowerLog{
Type: types.PowerConsume,
Model: generator.Provider(),
Remark: fmt.Sprintf("PPT 任务 %s 第 %d 页图片生成", task.TaskID, slide.SlideIndex),
})
if err != nil {
return fmt.Errorf("扣减算力失败:%v", err)
}
}
return nil
})
}
return group.Wait()
}
// startSlideImageGenerationAsync 在后台为 missing 页并发生图;若 setProcessingBeforeRun 为 true 则先置为 processing(用户主动 resume)。
func (s *PptService) startSlideImageGenerationAsync(task *Task, missing []SlideData, setProcessingBeforeRun bool) error {
if len(missing) == 0 {
return nil
}
cfg, err := s.loadPPTConfig()
if err != nil {
return err
}
cost := len(missing) * cfg.PowerCostPerSlide
if cost > 0 {
power, err := s.userPower(task.UserID)
if err != nil {
return err
}
if power < cost {
return ErrInsufficientPower
}
}
generator, err := NewImageGenerator(cfg)
if err != nil {
return fmt.Errorf("初始化图片生成器失败:%w", err)
}
if setProcessingBeforeRun {
s.UpdateStatus(task.TaskID, TaskStatusProcessing)
}
t := task
go func() {
bg := context.Background()
if err := s.runSlideImageJobs(bg, t, cfg, generator, missing); err != nil {
s.MarkAsFailed(task.TaskID, fmt.Sprintf("图片生成失败:%v", err))
return
}
s.UpdateStatus(task.TaskID, TaskStatusCompleted)
}()
return nil
}
// RecoverStaleProcessingTasks 进程启动时扫描 DB 中仍为 processing 且存在缺图页的任务,重新拉起生图协程(用于服务中断后的恢复)。
func (s *PptService) RecoverStaleProcessingTasks() {
var jobs []model.PPTJob
if err := s.db.Where("status = ?", string(TaskStatusProcessing)).Find(&jobs).Error; err != nil {
logger.Warnf("PPT recover: list processing jobs failed: %v", err)
return
}
for i := range jobs {
task := modelToTask(&jobs[i])
missing := slidesNeedingImages(task)
if len(missing) == 0 {
s.UpdateStatus(task.TaskID, TaskStatusCompleted)
logger.Infof("PPT recover: task %s was processing but all slides had images, marked completed", task.TaskID)
continue
}
if err := s.validateSlideOutline(task); err != nil {
logger.Warnf("PPT recover: task %s skip (invalid outline): %v", task.TaskID, err)
continue
}
if err := s.startSlideImageGenerationAsync(task, missing, false); err != nil {
if errors.Is(err, ErrInsufficientPower) {
logger.Warnf("PPT recover: task %s skip (insufficient power for %d slides)", task.TaskID, len(missing))
continue
}
logger.Warnf("PPT recover: task %s failed to restart: %v", task.TaskID, err)
continue
}
logger.Infof("PPT recover: restarted image generation for task %s (%d slides)", task.TaskID, len(missing))
}
}
// ResumeTask 继续为缺图页生图(需完整分镜占位;processing 时返回 ErrPptTaskBusy
func (s *PptService) ResumeTask(ctx context.Context, taskID string, userID uint) error {
task, ok := s.GetTask(taskID)
if !ok {
return ErrPptTaskNotFound
}
if task.UserID != userID {
return ErrPptTaskNotFound
}
if task.Status == TaskStatusProcessing {
return ErrPptTaskBusy
}
if task.Status == TaskStatusCompleted {
return ErrPptTaskNotResumable
}
if err := s.validateSlideOutline(task); err != nil {
return err
}
missing := slidesNeedingImages(task)
if len(missing) == 0 {
s.UpdateStatus(taskID, TaskStatusCompleted)
return nil
}
return s.startSlideImageGenerationAsync(task, missing, true)
}
// MarkAsFailed 标记任务失败
func (s *PptService) MarkAsFailed(taskID string, msg string) {
s.db.Model(&model.PPTJob{}).Where("task_id = ?", taskID).
Updates(map[string]interface{}{"status": string(TaskStatusFailed), "err_msg": msg, "updated_at": time.Now()})
}
// DeleteTask 删除任务并删除关联生成图片
// 仅允许删除 completed / failed 状态的任务,避免并发任务生成过程被打断。
func (s *PptService) DeleteTask(taskID string, userID uint) error {
var job model.PPTJob
if err := s.db.Where("task_id = ? AND user_id = ?", taskID, userID).First(&job).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) || job.TaskId == "" {
return ErrPptTaskNotFound
}
return err
}
if job.Status != string(TaskStatusCompleted) && job.Status != string(TaskStatusFailed) {
return ErrPptTaskNotDeletable
}
// 删除所有幻灯片对应的图片对象。
uploader := s.uploadManager.GetUploadHandler()
for _, slide := range job.Slides {
if slide.ImageURL == "" {
continue
}
if err := uploader.Delete(slide.ImageURL); err != nil {
// 图片可能已过期/不存在/对象已被清理,此时不应阻断“删除任务记录”的主流程。
// 这里只记录日志,确保数据库记录被删除后前端认为任务删除成功。
logger.Warnf("delete ppt image failed (task_id=%s, url=%s): %v", taskID, slide.ImageURL, err)
}
}
// 最后删除任务记录(slides 会随之从数据库消失)。
return s.db.Where("task_id = ? AND user_id = ?", taskID, userID).Delete(&model.PPTJob{}).Error
}
// List 返回所有任务列表(从数据库按创建时间倒序),供调用方按用户/状态过滤与分页
func (s *PptService) List() []*Task {
var jobs []model.PPTJob
s.db.Order("created_at DESC").Find(&jobs)
tasks := make([]*Task, 0, len(jobs))
for i := range jobs {
tasks = append(tasks, modelToTask(&jobs[i]))
}
sort.Slice(tasks, func(i, j int) bool {
return tasks[i].CreatedAt.After(tasks[j].CreatedAt)
})
return tasks
}
// ListUserTasks 当前用户的任务分页列表(缺 title/thumb 时补全并写库)
func (s *PptService) ListUserTasks(ctx context.Context, userID uint, page, pageSize int) ([]*Task, int) {
all := s.List()
filtered := make([]*Task, 0, len(all))
for _, t := range all {
if t.UserID == userID {
filtered = append(filtered, t)
}
}
total := len(filtered)
start := (page - 1) * pageSize
if start > total {
start = total
}
end := start + pageSize
if end > total {
end = total
}
slice := filtered[start:end]
cfg, cfgErr := s.loadPPTConfig()
if cfgErr != nil {
logger.Warnf("ListUserTasks loadPPTConfig: %v", cfgErr)
}
if s.llm == nil {
s.llm = NewLLMClient()
}
for _, t := range slice {
s.ensureTaskMeta(ctx, t, cfg)
}
return slice, total
}
// EnsureTaskMeta 对外暴露标题/缩略图补全逻辑,供管理端列表/详情复用。
func (s *PptService) EnsureTaskMeta(ctx context.Context, task *Task) {
if task == nil {
return
}
cfg, cfgErr := s.loadPPTConfig()
if cfgErr != nil {
logger.Warnf("EnsureTaskMeta loadPPTConfig: %v", cfgErr)
}
if s.llm == nil {
s.llm = NewLLMClient()
}
s.ensureTaskMeta(ctx, task, cfg)
}
func (s *PptService) ensureTaskMeta(ctx context.Context, task *Task, cfg types.PPTConfig) {
updates := map[string]interface{}{}
if strings.TrimSpace(task.Title) == "" && strings.TrimSpace(task.Content) != "" {
title := "未命名演示文稿"
if cfg.OutlineLLMApiURL != "" && cfg.OutlineLLMApiKey != "" {
ti, err := s.llm.GeneratePPTTitle(ctx, cfg, task.Content, task.Language)
if err != nil {
logger.Warnf("ensureTaskMeta GeneratePPTTitle task_id=%s: %v", task.TaskID, err)
} else {
ti = strings.TrimSpace(ti)
if ti != "" {
title = truncateTitleRunes(ti, 255)
}
}
}
task.Title = title
updates["title"] = task.Title
}
if task.Thumb == "" && len(task.Slides) > 0 {
thumb := DerivePPTThumbFromSlides(task.Slides)
if thumb != "" {
task.Thumb = thumb
updates["thumb"] = thumb
}
}
if len(updates) > 0 {
updates["updated_at"] = time.Now()
_ = s.db.Model(&model.PPTJob{}).Where("task_id = ?", task.TaskID).Updates(updates).Error
}
}
// ListAdminJobs 管理后台任务列表(筛选 + 分页)
func (s *PptService) ListAdminJobs(ctx context.Context, page, pageSize, filterUserID int, status string) ([]*Task, int) {
items := s.List()
filtered := make([]*Task, 0, len(items))
for _, t := range items {
if filterUserID > 0 && int(t.UserID) != filterUserID {
continue
}
if status != "" && string(t.Status) != status {
continue
}
filtered = append(filtered, t)
}
total := len(filtered)
if page <= 0 {
page = 1
}
if pageSize <= 0 {
pageSize = 20
}
start := (page - 1) * pageSize
if start > total {
start = total
}
end := start + pageSize
if end > total {
end = total
}
slice := filtered[start:end]
for _, t := range slice {
s.EnsureTaskMeta(ctx, t)
}
return slice, total
}
// Stats 任务状态统计(管理后台)
func (s *PptService) Stats() (total, completed, processing, failed, pending int64) {
for _, t := range s.List() {
total++
switch t.Status {
case TaskStatusCompleted:
completed++
case TaskStatusProcessing:
processing++
case TaskStatusFailed:
failed++
case TaskStatusPending:
pending++
}
}
return
}
// RunTask 执行 PPT 生成:分镜、并发生图、转存图片、扣算力
func (s *PptService) RunTask(ctx context.Context, task *Task, cfg types.PPTConfig) {
s.UpdateStatus(task.TaskID, TaskStatusProcessing)
maxPages := task.Pages
if maxPages <= 0 {
maxPages = cfg.MaxSlidesPerTask
}
plans, err := s.llm.GenerateSlides(ctx, cfg, task.Content, task.Prompt, task.Language, task.Mode, maxPages)
if err != nil {
s.MarkAsFailed(task.TaskID, fmt.Sprintf("生成分镜失败:%v", err))
return
}
if len(plans) == 0 {
s.MarkAsFailed(task.TaskID, "分镜结果为空")
return
}
total := len(plans)
if cfg.MaxSlidesPerTask > 0 && total > cfg.MaxSlidesPerTask {
plans = plans[:cfg.MaxSlidesPerTask]
total = len(plans)
}
outlines := slidePlansToOutlines(plans)
if err := s.saveSlidesOutline(task.TaskID, total, outlines); err != nil {
s.MarkAsFailed(task.TaskID, fmt.Sprintf("保存分镜占位失败:%v", err))
return
}
generator, err := NewImageGenerator(cfg)
if err != nil {
s.MarkAsFailed(task.TaskID, fmt.Sprintf("初始化图片生成器失败:%v", err))
return
}
if err := s.runSlideImageJobs(ctx, task, cfg, generator, outlines); err != nil {
s.MarkAsFailed(task.TaskID, fmt.Sprintf("图片生成失败:%v", err))
return
}
s.UpdateStatus(task.TaskID, TaskStatusCompleted)
}
func (s *PptService) loadPPTConfig() (types.PPTConfig, error) {
var cfgModel model.Config
var pptCfg types.PPTConfig
err := s.db.Where("name", types.ConfigKeyPPT).First(&cfgModel).Error
if err != nil {
if err == gorm.ErrRecordNotFound {
pptCfg.MaxSlidesPerTask = 30
pptCfg.MaxConcurrentRequests = 3
pptCfg.QPSLimit = 1
pptCfg.PowerCostPerSlide = 0
return pptCfg, nil
}
return pptCfg, err
}
err = utils.JsonDecode(cfgModel.Value, &pptCfg)
if err != nil {
return pptCfg, err
}
legacyMax10 := pptCfg.MaxSlidesPerTask == 10
if pptCfg.MaxSlidesPerTask <= 0 {
pptCfg.MaxSlidesPerTask = 30
}
if legacyMax10 {
// 与前端 PPT 页数控件 max=30 对齐;历史默认 10 会导致用户选择 12/15 仍被截断为 10
pptCfg.MaxSlidesPerTask = 30
}
if pptCfg.MaxConcurrentRequests <= 0 {
pptCfg.MaxConcurrentRequests = 3
}
if pptCfg.QPSLimit <= 0 {
pptCfg.QPSLimit = 1
}
if legacyMax10 {
val := utils.JsonEncode(pptCfg)
_ = s.db.Model(&model.Config{}).Where("name = ?", types.ConfigKeyPPT).Update("value", val)
}
return pptCfg, nil
}
+50
View File
@@ -0,0 +1,50 @@
package ppt
import (
"encoding/base64"
"fmt"
"geekai/core/types"
"net/http"
"net/url"
"os"
"strings"
)
// PrepareReferenceInputsForImg2Img 将幻灯片参考图转为第三方 API 可消费的输入:本地存储时读文件并转为 data URI(base64),公网 URL 原样传递。
func PrepareReferenceInputsForImg2Img(rawURL string, oss types.OSSConfig, app *types.AppConfig) ([]string, error) {
rawURL = strings.TrimSpace(rawURL)
if rawURL == "" {
return nil, fmt.Errorf("empty reference image url")
}
if strings.HasPrefix(rawURL, "data:") {
return []string{rawURL}, nil
}
if oss.Active == "local" {
if fp, ok := mapLocalUploadFile(rawURL, oss.Local); ok {
b, err := os.ReadFile(fp)
if err != nil {
return nil, fmt.Errorf("read local reference image: %w", err)
}
if _, err := decodeImageBytes(b); err != nil {
return nil, fmt.Errorf("reference is not a valid image: %w", err)
}
mime := http.DetectContentType(b)
if !strings.HasPrefix(mime, "image/") {
mime = "image/png"
}
dataURI := fmt.Sprintf("data:%s;base64,%s", mime, base64.StdEncoding.EncodeToString(b))
return []string{dataURI}, nil
}
}
if u, err := url.Parse(rawURL); err == nil && u.Scheme != "" && u.Host != "" {
return []string{rawURL}, nil
}
abs := resolveAbsoluteImageURL(rawURL, oss.Local, app)
if strings.TrimSpace(abs) == "" {
return nil, fmt.Errorf("cannot resolve reference image url")
}
return []string{abs}, nil
}
+182
View File
@@ -0,0 +1,182 @@
package ppt
import (
"context"
"fmt"
"geekai/core/types"
"geekai/store/model"
"geekai/store/vo"
"strings"
)
// EditSlideImage 基于当前激活图做图生图,追加 image_history 并将 image_url 设为新版。
func (s *PptService) EditSlideImage(ctx context.Context, taskID string, userID uint, slideIndex int, prompt string, oss types.OSSConfig, app *types.AppConfig) ([]SlideData, error) {
prompt = strings.TrimSpace(prompt)
if prompt == "" {
return nil, fmt.Errorf("请输入修改说明")
}
task, ok := s.GetTask(taskID)
if !ok {
return nil, ErrPptTaskNotFound
}
if task.UserID != userID {
return nil, ErrPptTaskNotFound
}
refURL := ""
for _, sl := range task.Slides {
if sl.SlideIndex == slideIndex {
normalizeSlideImageHistory(&sl)
refURL = strings.TrimSpace(sl.ImageURL)
break
}
}
if refURL == "" {
if slideExists(task.Slides, slideIndex) {
return nil, ErrPptSlideNoImage
}
return nil, ErrPptSlideNotFound
}
cfg, err := s.loadPPTConfig()
if err != nil {
return nil, err
}
power, err := s.userPower(userID)
if err != nil {
return nil, err
}
if cfg.PowerCostPerSlide > 0 && power < cfg.PowerCostPerSlide {
return nil, ErrInsufficientPower
}
generator, err := NewImageGenerator(cfg)
if err != nil {
return nil, err
}
refInputs, err := PrepareReferenceInputsForImg2Img(refURL, oss, app)
if err != nil {
return nil, fmt.Errorf("准备参考图失败:%w", err)
}
imgURL, err := generator.GenerateWithReference(ctx, prompt, refInputs)
if err != nil {
return nil, err
}
storedURL, err := s.uploadManager.GetUploadHandler().PutUrlFile(imgURL, ".png", false)
if err != nil {
return nil, fmt.Errorf("转存图片失败:%w", err)
}
if err := s.applySlideImageEdit(taskID, slideIndex, storedURL, prompt); err != nil {
return nil, err
}
if cfg.PowerCostPerSlide > 0 {
err = s.userService.DecreasePower(userID, cfg.PowerCostPerSlide, model.PowerLog{
Type: types.PowerConsume,
Model: generator.Provider(),
Remark: fmt.Sprintf("PPT 任务 %s 第 %d 页图生图编辑", taskID, slideIndex),
})
if err != nil {
return nil, fmt.Errorf("扣减算力失败:%v", err)
}
}
task2, _ := s.GetTask(taskID)
return task2.Slides, nil
}
func slideExists(slides []SlideData, slideIndex int) bool {
for _, sl := range slides {
if sl.SlideIndex == slideIndex {
return true
}
}
return false
}
func (s *PptService) applySlideImageEdit(taskID string, slideIndex int, newURL string, editPrompt string) error {
s.slidesLock.Lock()
defer s.slidesLock.Unlock()
var job model.PPTJob
if err := s.db.Where("task_id = ?", taskID).First(&job).Error; err != nil {
return err
}
slides := job.Slides
found := false
for i := range slides {
if slides[i].SlideIndex != slideIndex {
continue
}
found = true
sd := voToSlide(slides[i])
normalizeSlideImageHistory(&sd)
if strings.TrimSpace(sd.ImageURL) == "" {
return ErrPptSlideNoImage
}
sd.ImageHistory = append(sd.ImageHistory, vo.PPTSlideImageVersion{
ImageURL: newURL,
Prompt: editPrompt,
})
sd.ImageURL = newURL
slides[i] = slideToVO(sd)
break
}
if !found {
return ErrPptSlideNotFound
}
job.Slides = slides
biz := voSlidesToBiz(slides)
return s.refreshJobMeta(&job, biz)
}
// SetActiveSlideVersion 将 image_url 切换为 image_history[versionIndex]。
func (s *PptService) SetActiveSlideVersion(taskID string, userID uint, slideIndex int, versionIndex int) ([]SlideData, error) {
task, ok := s.GetTask(taskID)
if !ok {
return nil, ErrPptTaskNotFound
}
if task.UserID != userID {
return nil, ErrPptTaskNotFound
}
s.slidesLock.Lock()
defer s.slidesLock.Unlock()
var job model.PPTJob
if err := s.db.Where("task_id = ?", taskID).First(&job).Error; err != nil {
return nil, err
}
slides := job.Slides
found := false
for i := range slides {
if slides[i].SlideIndex != slideIndex {
continue
}
found = true
sd := voToSlide(slides[i])
normalizeSlideImageHistory(&sd)
hist := sd.ImageHistory
if versionIndex < 0 || versionIndex >= len(hist) {
return nil, ErrPptInvalidVersionIndex
}
sd.ImageURL = hist[versionIndex].ImageURL
slides[i] = slideToVO(sd)
break
}
if !found {
return nil, ErrPptSlideNotFound
}
job.Slides = slides
biz := voSlidesToBiz(slides)
if err := s.refreshJobMeta(&job, biz); err != nil {
return nil, err
}
out := voSlidesToBizNormalized(job.Slides)
return out, nil
}
-299
View File
@@ -1,299 +0,0 @@
package sd
// * +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
// * 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/types"
logger2 "geekai/logger"
"geekai/service"
"geekai/service/oss"
"geekai/store"
"geekai/store/model"
"geekai/utils"
"time"
"github.com/go-redis/redis/v8"
"github.com/imroc/req/v3"
"gorm.io/gorm"
)
var logger = logger2.GetLogger()
// SD 绘画服务
type Service struct {
httpClient *req.Client
taskQueue *store.RedisQueue
db *gorm.DB
uploadManager *oss.UploaderManager
userService *service.UserService
}
func NewService(db *gorm.DB, manager *oss.UploaderManager, redisCli *redis.Client, userService *service.UserService) *Service {
return &Service{
httpClient: req.C(),
taskQueue: store.NewRedisQueue("StableDiffusion_Task_Queue", redisCli),
db: db,
uploadManager: manager,
userService: userService,
}
}
func (s *Service) Run() {
// 将数据库中未提交的人物加载到队列
var jobs []model.SdJob
s.db.Where("progress", 0).Find(&jobs)
for _, v := range jobs {
var task types.SdTask
err := utils.JsonDecode(v.TaskInfo, &task)
if err != nil {
logger.Errorf("decode task info with error: %v", err)
continue
}
task.Id = int(v.Id)
s.PushTask(task)
}
logger.Infof("Starting Stable-Diffusion job consumer")
go func() {
for {
var task types.SdTask
err := s.taskQueue.LPop(&task)
if err != nil {
logger.Errorf("taking task with error: %v", err)
continue
}
// translate prompt
if utils.HasChinese(task.Params.Prompt) {
content, err := utils.OpenAIRequest(s.db, fmt.Sprintf(service.TranslatePromptTemplate, task.Params.Prompt), task.TranslateModelId)
if err == nil {
task.Params.Prompt = content
} else {
logger.Warnf("error with translate prompt: %v", err)
}
}
// translate negative prompt
if task.Params.NegPrompt != "" && utils.HasChinese(task.Params.NegPrompt) {
content, err := utils.OpenAIRequest(s.db, fmt.Sprintf(service.TranslatePromptTemplate, task.Params.NegPrompt), task.TranslateModelId)
if err == nil {
task.Params.NegPrompt = content
} else {
logger.Warnf("error with translate prompt: %v", err)
}
}
logger.Infof("handle a new Stable-Diffusion task: %+v", task)
err = s.Txt2Img(task)
if err != nil {
logger.Error("绘画任务执行失败:", err.Error())
// update the task progress
s.db.Model(&model.SdJob{Id: uint(task.Id)}).UpdateColumns(map[string]interface{}{
"progress": service.FailTaskProgress,
"err_msg": err.Error(),
})
continue
}
}
}()
}
// Txt2ImgReq 文生图请求实体
type Txt2ImgReq struct {
Prompt string `json:"prompt"`
NegativePrompt string `json:"negative_prompt"`
Seed int64 `json:"seed,omitempty"`
Steps int `json:"steps"`
CfgScale float32 `json:"cfg_scale"`
Width int `json:"width"`
Height int `json:"height"`
SamplerName string `json:"sampler_name"`
Scheduler string `json:"scheduler"`
EnableHr bool `json:"enable_hr,omitempty"`
HrScale int `json:"hr_scale,omitempty"`
HrUpscaler string `json:"hr_upscaler,omitempty"`
HrSecondPassSteps int `json:"hr_second_pass_steps,omitempty"`
DenoisingStrength float32 `json:"denoising_strength,omitempty"`
ForceTaskId string `json:"force_task_id,omitempty"`
}
// Txt2ImgResp 文生图响应实体
type Txt2ImgResp struct {
Images []string `json:"images"`
Parameters struct {
} `json:"parameters"`
Info string `json:"info"`
}
// TaskProgressResp 任务进度响应实体
type TaskProgressResp struct {
Progress float64 `json:"progress"`
EtaRelative float64 `json:"eta_relative"`
}
// Txt2Img 文生图 API
func (s *Service) Txt2Img(task types.SdTask) error {
body := Txt2ImgReq{
Prompt: task.Params.Prompt,
NegativePrompt: task.Params.NegPrompt,
Steps: task.Params.Steps,
CfgScale: task.Params.CfgScale,
Width: task.Params.Width,
Height: task.Params.Height,
SamplerName: task.Params.Sampler,
Scheduler: task.Params.Scheduler,
ForceTaskId: task.Params.TaskId,
}
if task.Params.Seed > 0 {
body.Seed = task.Params.Seed
}
if task.Params.HdFix {
body.EnableHr = true
body.HrScale = task.Params.HdScale
body.HrUpscaler = task.Params.HdScaleAlg
body.HrSecondPassSteps = task.Params.HdSteps
body.DenoisingStrength = task.Params.HdRedrawRate
}
var res Txt2ImgResp
var errChan = make(chan error)
var apiKey model.ApiKey
err := s.db.Where("type", "sd").Where("enabled", true).Order("last_used_at ASC").First(&apiKey).Error
if err != nil {
return fmt.Errorf("no available Stable-Diffusion api key: %v", err)
}
apiURL := fmt.Sprintf("%s/sdapi/v1/txt2img", apiKey.ApiURL)
logger.Infof("send image request to %s", apiURL)
// send a request to sd api endpoint
go func() {
response, err := s.httpClient.R().
SetHeader("Authorization", apiKey.Value).
SetBody(body).
SetSuccessResult(&res).
Post(apiURL)
if err != nil {
errChan <- err
return
}
if response.IsErrorState() {
errChan <- fmt.Errorf("error http code status: %v", response.Status)
return
}
// update the last used time
apiKey.LastUsedAt = time.Now().Unix()
s.db.Updates(&apiKey)
// 保存 Base64 图片
imgURL, err := s.uploadManager.GetUploadHandler().PutBase64(res.Images[0])
if err != nil {
errChan <- fmt.Errorf("error with upload image: %v", err)
return
}
// 获取绘画真实的 seed
var info map[string]interface{}
err = utils.JsonDecode(res.Info, &info)
if err != nil {
errChan <- fmt.Errorf("error with decode task response: %v", err)
return
}
task.Params.Seed = int64(utils.IntValue(utils.InterfaceToString(info["seed"]), -1))
s.db.Model(&model.SdJob{Id: uint(task.Id)}).UpdateColumns(model.SdJob{ImgURL: imgURL, Params: utils.JsonEncode(task.Params), Prompt: task.Params.Prompt})
errChan <- nil
}()
// waiting for task finish
for {
select {
case err := <-errChan:
if err != nil {
return err
}
// task finished
s.db.Model(&model.SdJob{Id: uint(task.Id)}).UpdateColumn("progress", 100)
return nil
default:
resp, err := s.checkTaskProgress(apiKey)
// 更新任务进度
if err == nil && resp.Progress > 0 {
s.db.Model(&model.SdJob{Id: uint(task.Id)}).UpdateColumn("progress", int(resp.Progress*100))
}
time.Sleep(time.Second)
}
}
}
// 执行任务
func (s *Service) checkTaskProgress(apiKey model.ApiKey) (*TaskProgressResp, error) {
apiURL := fmt.Sprintf("%s/sdapi/v1/progress?skip_current_image=false", apiKey.ApiURL)
var res TaskProgressResp
response, err := s.httpClient.R().
SetHeader("Authorization", apiKey.Value).
SetSuccessResult(&res).
Get(apiURL)
if err != nil {
return nil, err
}
if response.IsErrorState() {
return nil, fmt.Errorf("error http code status: %v", response.Status)
}
return &res, nil
}
func (s *Service) PushTask(task types.SdTask) {
logger.Debugf("add a new MidJourney task to the task list: %+v", task)
if err := s.taskQueue.RPush(task); err != nil {
logger.Errorf("push sd task to queue failed: %v", err)
}
}
// CheckTaskStatus 检查任务状态,自动删除过期或者失败的任务
func (s *Service) CheckTaskStatus() {
go func() {
logger.Info("Running Stable-Diffusion task status checking ...")
for {
var jobs []model.SdJob
res := s.db.Where("progress < ?", 100).Find(&jobs)
if res.Error != nil {
time.Sleep(5 * time.Second)
continue
}
for _, job := range jobs {
// 5 分钟还没完成的任务标记为失败
if time.Since(job.CreatedAt) > time.Minute*5 {
job.Progress = service.FailTaskProgress
job.ErrMsg = "任务超时"
s.db.Updates(&job)
}
}
// 找出失败的任务,并恢复其扣减算力
s.db.Where("progress", service.FailTaskProgress).Where("power > ?", 0).Find(&jobs)
for _, job := range jobs {
err := s.userService.IncreasePower(job.UserId, job.Power, model.PowerLog{
Type: types.PowerRefund,
Model: "stable-diffusion",
Remark: fmt.Sprintf("任务失败,退回算力。任务ID%d Err: %s", job.Id, job.ErrMsg),
})
if err != nil {
continue
}
// 更新任务状态
s.db.Model(&job).UpdateColumn("power", 0)
}
time.Sleep(time.Second * 5)
}
}()
}
+4 -10
View File
@@ -8,14 +8,14 @@ package sms
// * +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
import (
"context"
"fmt"
"geekai/core/types"
"geekai/utils"
"io"
"net/http"
"net/url"
"strconv"
"strings"
"time"
)
type BaoSmsService struct {
@@ -58,15 +58,9 @@ func (s *BaoSmsService) SendVerifyCode(mobile string, code int) error {
params.Set("c", content)
apiURL := fmt.Sprintf("https://%s/sms?%s", s.domain, params.Encode())
response, err := http.Get(apiURL)
body, status, err := utils.FetchURLBytes(context.Background(), apiURL, "", 30*time.Second, 2, 2<<20)
if err != nil {
return err
}
defer response.Body.Close()
body, err := io.ReadAll(response.Body)
if err != nil {
return err
return fmt.Errorf("smsbao request failed: status=%d: %w", status, err)
}
result := string(body)
logger.Debugf("send SmsBao result: %v", errMsg[result])
+1
View File
@@ -9,6 +9,7 @@ package sms
const Ali = "aliyun"
const Bao = "bao"
const Tencent = "tencent"
type Service interface {
SendVerifyCode(mobile string, code int) error
+15 -9
View File
@@ -9,23 +9,25 @@ package sms
import (
"geekai/core/types"
logger2 "geekai/logger"
"geekai/log"
)
type SmsManager struct {
aliyun *AliYunSmsService
bao *BaoSmsService
active string
aliyun *AliYunSmsService
bao *BaoSmsService
tencent *TencentSmsService
active string
}
var logger = logger2.GetLogger()
var logger = log.GetLogger()
func NewSmsManager(sysConfig *types.SystemConfig, aliyun *AliYunSmsService, bao *BaoSmsService) (*SmsManager, error) {
func NewSmsManager(sysConfig *types.SystemConfig, aliyun *AliYunSmsService, bao *BaoSmsService, tencent *TencentSmsService) (*SmsManager, error) {
return &SmsManager{
active: sysConfig.SMS.Active,
aliyun: aliyun,
bao: bao,
active: sysConfig.SMS.Active,
aliyun: aliyun,
bao: bao,
tencent: tencent,
}, nil
}
@@ -35,6 +37,8 @@ func (m *SmsManager) GetService() Service {
return m.aliyun
case Bao:
return m.bao
case Tencent:
return m.tencent
}
return nil
}
@@ -49,6 +53,8 @@ func (m *SmsManager) UpdateConfig(config types.SMSConfig) {
m.aliyun.UpdateConfig(config.Ali)
case Bao:
m.bao.UpdateConfig(config.Bao)
case Tencent:
m.tencent.UpdateConfig(config.Tencent)
}
m.active = config.Active
}
+127
View File
@@ -0,0 +1,127 @@
package sms
// * +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
// * 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/types"
"github.com/tencentcloud/tencentcloud-sdk-go/tencentcloud/common"
"github.com/tencentcloud/tencentcloud-sdk-go/tencentcloud/common/profile"
sms "github.com/tencentcloud/tencentcloud-sdk-go/tencentcloud/sms/v20210111"
)
type TencentSmsService struct {
config types.SmsConfigTencent
client *sms.Client
region string
}
func NewTencentSmsService(sysConfig *types.SystemConfig) (*TencentSmsService, error) {
config := sysConfig.SMS.Tencent
region := config.Region
if region == "" {
region = "ap-guangzhou" // 默认使用广州地区
}
s := TencentSmsService{
config: config,
region: region,
}
if sysConfig.SMS.Active == Tencent {
err := s.UpdateConfig(config)
if err != nil {
logger.Errorf("腾讯云短信初始化失败: %v", err)
}
}
return &s, nil
}
func (s *TencentSmsService) UpdateConfig(config types.SmsConfigTencent) error {
if config.SecretId == "" || config.SecretKey == "" {
// 配置不完整时不初始化客户端
s.config = config
if config.Region != "" {
s.region = config.Region
} else {
s.region = "ap-guangzhou"
}
return nil
}
region := config.Region
if region == "" {
region = "ap-guangzhou"
}
// 创建凭证
credential := common.NewCredential(
config.SecretId,
config.SecretKey,
)
// 创建客户端配置
cpf := profile.NewClientProfile()
cpf.HttpProfile.Endpoint = "sms.tencentcloudapi.com"
// 创建客户端
client, err := sms.NewClient(credential, region, cpf)
if err != nil {
return fmt.Errorf("failed to create client: %v", err)
}
s.client = client
s.config = config
s.region = region
return nil
}
// SendVerifyCode 发送验证码短信
// 注意:腾讯云后台配置的短信模板内容应与配置中的 code_template 一致
// 模板只需要1个参数:{1} 表示验证码,例如:{1}为您的验证码,请于5分钟内填写,如非本人操作,请忽略本短信。
func (s *TencentSmsService) SendVerifyCode(mobile string, code int) error {
if s.client == nil {
return fmt.Errorf("腾讯云短信服务未初始化")
}
// 创建发送短信请求
request := sms.NewSendSmsRequest()
request.SmsSdkAppId = common.StringPtr(s.config.SmsSdkAppId)
request.SignName = common.StringPtr(s.config.Sign)
request.TemplateId = common.StringPtr(s.config.CodeTempId)
request.PhoneNumberSet = common.StringPtrs([]string{mobile})
request.TemplateParamSet = common.StringPtrs([]string{fmt.Sprintf("%d", code), "5"})
// 发送短信
response, err := s.client.SendSms(request)
if err != nil {
return fmt.Errorf("failed to send SMS: %v", err)
}
// 检查响应
if response.Response == nil {
return fmt.Errorf("failed to send SMS: response is nil")
}
if len(response.Response.SendStatusSet) == 0 {
return fmt.Errorf("failed to send SMS: no send status")
}
sendStatus := response.Response.SendStatusSet[0]
if sendStatus.Code == nil || *sendStatus.Code != "Ok" {
message := "unknown error"
if sendStatus.Message != nil {
message = *sendStatus.Message
}
return fmt.Errorf("failed to send SMS: %s", message)
}
return nil
}
var _ Service = &TencentSmsService{}
+1 -1
View File
@@ -120,7 +120,7 @@ func (s *SmtpService) sendTLS(auth smtp.Auth, to string, subject string, body st
}
_, _ = fmt.Fprintln(wc)
// 将邮件内容写入
_, err = fmt.Fprintf(wc, body)
_, err = fmt.Fprint(wc, body)
if err != nil {
return fmt.Errorf("error sending email: %v", err)
}
+8 -13
View File
@@ -1,21 +1,20 @@
package sora
import (
"context"
"encoding/json"
"errors"
"geekai/service/oss"
"geekai/store/vo"
"geekai/utils"
"io"
"net/http"
"path/filepath"
"regexp"
"time"
logger2 "geekai/logger"
"geekai/log"
)
var logger = logger2.GetLogger()
var logger = log.GetLogger()
type SoraService struct {
uploadManager *oss.UploaderManager
@@ -34,15 +33,10 @@ func (s *SoraService) DownloadVideoURL(text string) (*vo.File, error) {
return nil, err
}
// 获取 JSON 数据
resp, err := http.Get(videoDataURL)
if err != nil {
return nil, err
}
defer resp.Body.Close()
body, err := io.ReadAll(resp.Body)
// 用统一的超时/重试策略,避免“偶发 HTTPS 握手超时”直接导致失败
body, _, err := utils.FetchURLBytes(context.Background(), videoDataURL, "", 30*time.Second, 2, 2<<20)
if err != nil {
logger.Errorf("failed to get video data: %v", err)
return nil, err
}
@@ -50,13 +44,14 @@ func (s *SoraService) DownloadVideoURL(text string) (*vo.File, error) {
var videoData map[string]any
err = json.Unmarshal(body, &videoData)
if err != nil {
logger.Errorf("failed to unmarshal video data: %v", err)
return nil, err
}
if v, ok := videoData["url"].(string); ok && v != "" {
logger.Infof("try to download video: %s", v)
videoURL, err := s.uploadManager.GetUploadHandler().PutUrlFile(v, ".mp4", true)
if err != nil {
if err != nil { // 如果上传失败,则返回原始错误
return nil, err
}
+36 -16
View File
@@ -12,11 +12,12 @@ import (
"errors"
"fmt"
"geekai/core/types"
logger2 "geekai/logger"
"geekai/log"
"geekai/service"
"geekai/service/oss"
"geekai/store"
"geekai/store/model"
"geekai/store/vo"
"geekai/utils"
"io"
"time"
@@ -27,7 +28,7 @@ import (
"gorm.io/gorm"
)
var logger = logger2.GetLogger()
var logger = log.GetLogger()
type Service struct {
httpClient *req.Client
@@ -61,13 +62,24 @@ func (s *Service) Run() {
var jobs []model.SunoJob
s.db.Where("task_id", "").Where("progress", 0).Find(&jobs)
for _, v := range jobs {
var task types.SunoTask
err := utils.JsonDecode(v.TaskInfo, &task)
if err != nil {
logger.Errorf("decode task info with error: %v", err)
continue
// 从 Params 中提取字段构建 task
task := types.SunoTask{
Id: v.Id,
UserId: int(v.UserId),
Channel: v.Channel,
Type: v.Type,
Title: v.Title,
RefTaskId: v.RefTaskId,
RefSongId: v.RefSongId,
Prompt: v.Params.Prompt,
Lyrics: v.Params.Lyrics,
Tags: v.Params.Tags,
Model: v.Params.Model,
Instrumental: v.Params.Instrumental,
ExtendSecs: v.Params.ExtendSecs,
SongId: v.SongId,
AudioURL: v.AudioURL,
}
task.Id = v.Id
s.PushTask(task)
}
logger.Info("Starting Suno job consumer...")
@@ -335,15 +347,22 @@ func (s *Service) SyncTaskProgress() {
job.SongId = v.Id
job.Duration = int(v.Metadata.Duration)
job.Prompt = v.Metadata.Prompt
// 设置 Params
tags := v.Metadata.Tags
// 修复 tags 字段过长导致插入数据库失败
if len(v.Metadata.Tags) > 255 {
job.Tags = v.Metadata.Tags[:255]
} else {
job.Tags = v.Metadata.Tags
if len(tags) > 255 {
tags = tags[:255]
}
job.Params = vo.SunoParam{
Prompt: v.Metadata.Prompt,
Tags: tags,
Model: v.ModelName,
Instrumental: job.Params.Instrumental, // 保持原任务参数
ExtendSecs: job.Params.ExtendSecs, // 保持原任务参数
}
job.ModelName = v.ModelName
job.RawData = utils.JsonEncode(v)
job.Output = utils.JsonEncode(v)
job.CoverURL = v.ImageLargeUrl
job.AudioURL = v.AudioUrl
@@ -372,11 +391,12 @@ func (s *Service) SyncTaskProgress() {
}
// 找出失败的任务,并恢复其扣减算力
s.db.Where("progress", service.FailTaskProgress).Where("power > ?", 0).Find(&jobs)
s.db.Select("id", "user_id", "power", "task_id", "err_msg", "params").
Where("progress", service.FailTaskProgress).Where("power > ?", 0).Find(&jobs)
for _, job := range jobs {
err := s.userService.IncreasePower(job.UserId, job.Power, model.PowerLog{
Type: types.PowerRefund,
Model: job.ModelName,
Model: job.Params.Model,
Remark: fmt.Sprintf("Suno 任务失败,退回算力。任务ID%sErr:%s", job.TaskId, job.ErrMsg),
})
if err != nil {
+2 -2
View File
@@ -1,6 +1,6 @@
package service
import logger2 "geekai/logger"
import "geekai/log"
const FailTaskProgress = 101
const (
@@ -17,7 +17,7 @@ type NotifyMessage struct {
Type string `json:"type"`
}
var logger = logger2.GetLogger()
var logger = log.GetLogger()
const TranslatePromptTemplate = "Translate the following painting prompt words into English keyword phrases. Without any explanation, directly output the keyword phrases separated by commas. The content to be translated is: [%s]"
+32
View File
@@ -0,0 +1,32 @@
package adapters
import "geekai/log"
// * +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
// * 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
// * +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
var logger = log.GetLogger()
// CreateTaskResponse 创建任务响应
type CreateTaskResponse struct {
TaskId string `json:"task_id"` // 任务ID
Channel string `json:"channel"` // 渠道标识
Prompt string `json:"prompt"` // 优化后的提示词(如果有)
State string `json:"state"` // 任务状态
CreatedAt string `json:"created_at"` // 创建时间
}
// QueryTaskResponse 查询任务响应
type QueryTaskResponse struct {
TaskId string `json:"task_id"` // 任务ID
Status string `json:"status"` // 任务状态
Progress int `json:"progress"` // 进度(0-100
VideoURL string `json:"video_url"` // 视频URL
Prompt string `json:"prompt"` // 提示词
ErrMsg string `json:"err_msg"` // 错误信息
StatusMsg string `json:"status_msg"` // 状态消息
Output string `json:"output"` // 任务输出的原始信息(JSON字符串)
}
@@ -0,0 +1,302 @@
package adapters
// * +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
// * 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 (
"bytes"
"encoding/json"
"errors"
"fmt"
"geekai/core/types"
"io"
"net/http"
"time"
"gorm.io/gorm"
)
// DoubaoAdapter 豆包 Seedance 视频生成适配器(通过 Kapon VolcArk 接入)
type DoubaoAdapter struct {
db *gorm.DB
}
// NewDoubaoAdapter 创建 Doubao 适配器
func NewDoubaoAdapter(db *gorm.DB) *DoubaoAdapter {
return &DoubaoAdapter{
db: db,
}
}
// GetProvider 获取服务提供商名称
func (a *DoubaoAdapter) GetProvider() string {
return types.VideoDoubao
}
// doubaoContentItem 请求体中的 content 子项
type doubaoContentItem struct {
Type string `json:"type"`
Text string `json:"text,omitempty"`
ImageURL map[string]string `json:"image_url,omitempty"`
Extra map[string]interface{} `json:"extra,omitempty"` // 预留扩展
}
// doubaoCreateRequest 创建任务请求体
type doubaoCreateRequest struct {
Model string `json:"model"`
Content []doubaoContentItem `json:"content"`
Duration int `json:"duration,omitempty"`
Frames int `json:"frames,omitempty"`
Ratio string `json:"ratio,omitempty"`
Resolution string `json:"resolution,omitempty"`
Seed int64 `json:"seed,omitempty"`
// 其他官方支持的字段,按需追加
}
// doubaoCreateResponse 创建任务响应
type doubaoCreateResponse struct {
Id string `json:"id"`
PlatformId string `json:"platform_id"`
// 其余字段目前用不到,先不展开
}
// doubaoQueryContent 查询任务 content 字段
type doubaoQueryContent struct {
VideoURL string `json:"video_url"`
LastFrameURL string `json:"last_frame_url"`
Width int `json:"width"`
Height int `json:"height"`
}
// doubaoQueryUsage 查询任务 usage 字段
type doubaoQueryUsage struct {
VideoTokens int `json:"video_tokens"`
}
// doubaoQueryResponse 查询任务响应
type doubaoQueryResponse struct {
Id string `json:"id"`
PlatformId string `json:"platform_id"`
Model string `json:"model"`
Status string `json:"status"`
Content doubaoQueryContent `json:"content"`
Duration int `json:"duration"`
Frames int `json:"framespersecond"`
Usage doubaoQueryUsage `json:"usage"`
Error string `json:"error,omitempty"`
}
// CreateTask 创建豆包 Seedance 视频任务
func (a *DoubaoAdapter) CreateTask(task types.VideoTask, videoConfig *types.VideoConfig) (CreateTaskResponse, error) {
if videoConfig == nil {
return CreateTaskResponse{}, errors.New("视频配置为空")
}
if videoConfig.ApiURL == "" || videoConfig.ApiKey == "" {
return CreateTaskResponse{}, errors.New("豆包视频未配置 ApiURL 或 ApiKey")
}
paramsMap, ok := task.Params.(map[string]interface{})
if !ok {
return CreateTaskResponse{}, errors.New("invalid params type for Doubao video task")
}
// 模型名称:优先从 params.model 读取,否则默认 doubao-seedance-1-5-pro
modelName := "doubao-seedance-1-5-pro"
if v, ok := paramsMap["model"].(string); ok && v != "" {
modelName = v
}
// 解析基础参数
duration := 0
if v, ok := paramsMap["duration"].(float64); ok {
duration = int(v)
}
if v, ok := paramsMap["duration"].(int); ok {
duration = v
}
ratio := ""
if v, ok := paramsMap["aspect_ratio"].(string); ok {
ratio = v
}
resolution := ""
if v, ok := paramsMap["resolution"].(string); ok {
resolution = v
}
var seed int64
switch v := paramsMap["seed"].(type) {
case float64:
seed = int64(v)
case int:
seed = int64(v)
case int64:
seed = v
}
// 构建 content 数组:文本提示词为必填
content := []doubaoContentItem{
{
Type: "text",
Text: task.Prompt,
},
}
// 如果存在 input_reference(图片 URL),则追加 image_url 项,用于 I2V
if ref, ok := paramsMap["input_reference"].(string); ok && ref != "" {
content = append(content, doubaoContentItem{
Type: "image_url",
ImageURL: map[string]string{
"url": ref,
},
})
}
reqBody := doubaoCreateRequest{
Model: modelName,
Content: content,
Duration: duration,
Ratio: ratio,
Resolution: resolution,
Seed: seed,
}
payload, err := json.Marshal(reqBody)
if err != nil {
return CreateTaskResponse{}, fmt.Errorf("序列化豆包请求失败: %v", err)
}
logger.Debugf("DoubaoCreateRequest: %s", string(payload))
url := fmt.Sprintf("%s/seedance/v3/contents/generations/tasks", videoConfig.ApiURL)
req, err := http.NewRequest("POST", url, bytes.NewReader(payload))
if err != nil {
return CreateTaskResponse{}, fmt.Errorf("创建豆包请求失败: %v", err)
}
req.Header.Set("Authorization", "Bearer "+videoConfig.ApiKey)
req.Header.Set("Content-Type", "application/json")
client := &http.Client{Timeout: 60 * time.Second}
resp, err := client.Do(req)
if err != nil {
return CreateTaskResponse{}, fmt.Errorf("调用豆包接口失败: %v", err)
}
defer resp.Body.Close()
body, err := io.ReadAll(resp.Body)
if err != nil {
return CreateTaskResponse{}, fmt.Errorf("读取豆包响应失败: %v", err)
}
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
return CreateTaskResponse{}, fmt.Errorf("豆包接口返回错误状态码: %d, %s", resp.StatusCode, string(body))
}
var apiResp doubaoCreateResponse
if err := json.Unmarshal(body, &apiResp); err != nil {
return CreateTaskResponse{}, fmt.Errorf("解析豆包创建任务响应失败: %v, body=%s", err, string(body))
}
taskId := apiResp.PlatformId
if taskId == "" {
taskId = apiResp.Id
}
if taskId == "" {
return CreateTaskResponse{}, fmt.Errorf("豆包创建任务响应缺少任务 ID, body=%s", string(body))
}
return CreateTaskResponse{
TaskId: taskId,
Channel: videoConfig.ApiURL,
Prompt: task.Prompt,
State: types.VideoStatusPending,
CreatedAt: time.Now().Format(time.RFC3339),
}, nil
}
// QueryTask 查询豆包 Seedance 视频任务状态
func (a *DoubaoAdapter) QueryTask(taskId string, channel string, videoConfig *types.VideoConfig) (QueryTaskResponse, error) {
if videoConfig == nil {
return QueryTaskResponse{}, errors.New("视频配置为空")
}
if videoConfig.ApiURL == "" || videoConfig.ApiKey == "" {
return QueryTaskResponse{}, errors.New("豆包视频未配置 ApiURL 或 ApiKey")
}
url := fmt.Sprintf("%s/seedance/v3/contents/generations/tasks/%s", videoConfig.ApiURL, taskId)
req, err := http.NewRequest("GET", url, nil)
if err != nil {
return QueryTaskResponse{}, fmt.Errorf("创建查询请求失败: %v", err)
}
req.Header.Set("Authorization", "Bearer "+videoConfig.ApiKey)
req.Header.Set("Content-Type", "application/json")
client := &http.Client{Timeout: 60 * time.Second}
resp, err := client.Do(req)
if err != nil {
return QueryTaskResponse{}, fmt.Errorf("调用豆包查询接口失败: %v", err)
}
defer resp.Body.Close()
body, err := io.ReadAll(resp.Body)
if err != nil {
return QueryTaskResponse{}, fmt.Errorf("读取豆包查询响应失败: %v", err)
}
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
return QueryTaskResponse{}, fmt.Errorf("豆包查询接口返回错误状态码: %d, %s", resp.StatusCode, string(body))
}
var apiResp doubaoQueryResponse
if err := json.Unmarshal(body, &apiResp); err != nil {
return QueryTaskResponse{}, fmt.Errorf("解析豆包查询任务响应失败: %v, body=%s", err, string(body))
}
status := apiResp.Status
progress := 0
switch status {
case "queued":
status = types.VideoStatusPending
progress = 10
case "running":
status = types.VideoStatusInProgress
progress = 60
case "succeeded":
status = types.VideoStatusSuccess
progress = 100
case "failed", "cancelled":
status = types.VideoStatusFailed
default:
// 保持原样或视为 pending
status = types.VideoStatusPending
}
errMsg := apiResp.Error
if errMsg == "" && status == types.VideoStatusFailed {
errMsg = "doubao task failed"
}
result := QueryTaskResponse{
TaskId: apiResp.PlatformId,
Status: status,
Progress: progress,
VideoURL: apiResp.Content.VideoURL,
Prompt: "",
ErrMsg: errMsg,
StatusMsg: status,
Output: string(body),
}
if result.TaskId == "" {
result.TaskId = apiResp.Id
}
return result, nil
}
@@ -0,0 +1,267 @@
package adapters
// * +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
// * 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 (
"bytes"
"encoding/json"
"errors"
"fmt"
"geekai/core/types"
"io"
"net/http"
"time"
"gorm.io/gorm"
)
// KelingAdapter 可灵视频生成适配器
type KelingAdapter struct {
db *gorm.DB
}
// NewKelingAdapter 创建可灵适配器
func NewKelingAdapter(db *gorm.DB) *KelingAdapter {
return &KelingAdapter{
db: db,
}
}
// GetProvider 获取服务提供商名称
func (a *KelingAdapter) GetProvider() string {
return "keling"
}
// KelingCreateRequest 可灵创建任务请求
type KelingCreateRequest struct {
ModelName string `json:"model_name"`
Prompt string `json:"prompt"`
NegativePrompt string `json:"negative_prompt,omitempty"`
CfgScale float64 `json:"cfg_scale,omitempty"`
Mode string `json:"mode,omitempty"`
AspectRatio string `json:"aspect_ratio,omitempty"`
Duration string `json:"duration,omitempty"`
Sound bool `json:"sound,omitempty"`
Image string `json:"image,omitempty"`
ImageTail string `json:"image_tail,omitempty"`
}
// KelingCreateResponse 可灵创建任务响应
type KelingCreateResponse struct {
Code int `json:"code"`
Message string `json:"message"`
RequestID string `json:"request_id"`
Data struct {
TaskID string `json:"task_id"`
TaskStatus string `json:"task_status"`
CreatedAt int64 `json:"created_at"`
UpdatedAt int64 `json:"updated_at"`
} `json:"data"`
}
// KelingQueryResponse 可灵查询任务响应
type KelingQueryResponse struct {
Code int `json:"code"`
Message string `json:"message"`
Data struct {
TaskID string `json:"task_id"`
TaskStatus string `json:"task_status"`
TaskStatusMsg string `json:"task_status_msg"`
CreatedAt int64 `json:"created_at"`
UpdatedAt int64 `json:"updated_at"`
TaskResult struct {
Images []struct {
Index int `json:"index"`
URL string `json:"url"`
} `json:"images,omitempty"`
Videos []struct {
ID string `json:"id"`
URL string `json:"url"`
Duration string `json:"duration"`
} `json:"videos,omitempty"`
} `json:"task_result"`
} `json:"data"`
}
// CreateTask 创建视频生成任务
func (a *KelingAdapter) CreateTask(task types.VideoTask, videoConfig *types.VideoConfig) (CreateTaskResponse, error) {
// 解析任务参数
paramsMap, ok := task.Params.(map[string]interface{})
if !ok {
return CreateTaskResponse{}, errors.New("invalid params type for KeLing video task")
}
// 构建请求参数
payload := KelingCreateRequest{
Prompt: task.Prompt,
}
// 从 params 中提取参数
if modelName, ok := paramsMap["model_name"].(string); ok {
payload.ModelName = modelName
}
if prompt, ok := paramsMap["prompt"].(string); ok {
payload.Prompt = prompt
}
if negativePrompt, ok := paramsMap["negative_prompt"].(string); ok {
payload.NegativePrompt = negativePrompt
}
if cfgScale, ok := paramsMap["cfg_scale"].(float64); ok {
payload.CfgScale = cfgScale
}
if mode, ok := paramsMap["mode"].(string); ok {
payload.Mode = mode
}
if aspectRatio, ok := paramsMap["aspect_ratio"].(string); ok {
payload.AspectRatio = aspectRatio
}
if duration, ok := paramsMap["duration"].(string); ok {
payload.Duration = duration
}
if sound, ok := paramsMap["sound"].(bool); ok {
payload.Sound = sound
}
// 处理图生视频
taskType, ok := paramsMap["task_type"].(string)
if ok && taskType == "image2video" {
if image, ok := paramsMap["image"].(string); ok {
payload.Image = image
}
if imageTail, ok := paramsMap["image_tail"].(string); ok {
payload.ImageTail = imageTail
}
}
jsonPayload, err := json.Marshal(payload)
if err != nil {
return CreateTaskResponse{}, fmt.Errorf("failed to marshal payload: %v", err)
}
logger.Debugf("KelingCreateRequest: %+v", string(jsonPayload))
// 发送请求
url := fmt.Sprintf("%s/kling/v1/videos/%s", videoConfig.ApiURL, taskType)
req, err := http.NewRequest("POST", url, bytes.NewReader(jsonPayload))
if err != nil {
return CreateTaskResponse{}, fmt.Errorf("failed to create request: %v", err)
}
req.Header.Set("Authorization", "Bearer "+videoConfig.ApiKey)
req.Header.Set("Content-Type", "application/json")
// 发送请求
client := &http.Client{Timeout: time.Duration(30) * time.Second}
resp, err := client.Do(req)
if err != nil {
return CreateTaskResponse{}, fmt.Errorf("failed to send request: %v", err)
}
defer resp.Body.Close()
// 处理响应
body, err := io.ReadAll(resp.Body)
if err != nil {
return CreateTaskResponse{}, fmt.Errorf("failed to read response: %v", err)
}
if resp.StatusCode != http.StatusOK {
return CreateTaskResponse{}, fmt.Errorf("API error (status %d): %s", resp.StatusCode, string(body))
}
var apiResponse KelingCreateResponse
if err := json.Unmarshal(body, &apiResponse); err != nil {
return CreateTaskResponse{}, fmt.Errorf("failed to parse response: %v", err)
}
if apiResponse.Code != 0 {
return CreateTaskResponse{}, fmt.Errorf("API error: %s", apiResponse.Message)
}
return CreateTaskResponse{
TaskId: apiResponse.Data.TaskID,
Channel: videoConfig.ApiURL,
Prompt: task.Prompt,
State: types.VideoStatusPending,
CreatedAt: time.Now().Format(time.RFC3339),
}, nil
}
// QueryTask 查询任务状态
func (a *KelingAdapter) QueryTask(taskId string, channel string, videoConfig *types.VideoConfig) (QueryTaskResponse, error) {
// 从 taskId 中提取 action(可灵的 taskId 格式可能包含 action 信息)
// 这里需要从任务信息中获取 task_type,暂时使用 text2video 作为默认值
action := "text2video"
// 尝试从 channel 或其他地方获取 action,这里简化处理
// 实际应该从任务信息中获取
url := fmt.Sprintf("%s/kling/v1/videos/%s/%s", videoConfig.ApiURL, action, taskId)
req, err := http.NewRequest("GET", url, nil)
if err != nil {
return QueryTaskResponse{}, fmt.Errorf("failed to create request: %w", err)
}
req.Header.Set("Authorization", "Bearer "+videoConfig.ApiKey)
req.Header.Set("Content-Type", "application/json")
client := &http.Client{Timeout: time.Duration(30) * time.Second}
res, err := client.Do(req)
if err != nil {
return QueryTaskResponse{}, fmt.Errorf("failed to execute request: %w", err)
}
defer res.Body.Close()
if res.StatusCode != http.StatusOK {
body, _ := io.ReadAll(res.Body)
return QueryTaskResponse{}, fmt.Errorf("unexpected status code: %d, %s", res.StatusCode, string(body))
}
body, err := io.ReadAll(res.Body)
if err != nil {
return QueryTaskResponse{}, fmt.Errorf("failed to read response body: %w", err)
}
var response KelingQueryResponse
if err := json.Unmarshal(body, &response); err != nil {
return QueryTaskResponse{}, fmt.Errorf("failed to unmarshal response: %w", err)
}
if response.Code != 0 {
return QueryTaskResponse{}, fmt.Errorf("API error: %s", response.Message)
}
// 转换状态
state := response.Data.TaskStatus
status := state
switch state {
case "in_progress", "processing":
status = types.VideoStatusInProgress
case "completed", "succeed", "success":
status = types.VideoStatusSuccess
case "failed":
status = types.VideoStatusFailed
default:
status = types.VideoStatusPending
}
// 构建响应
result := QueryTaskResponse{
TaskId: response.Data.TaskID,
Status: status,
ErrMsg: response.Data.TaskStatusMsg,
StatusMsg: response.Data.TaskStatusMsg,
Output: string(body),
}
// 提取视频URL
if len(response.Data.TaskResult.Videos) > 0 {
result.VideoURL = response.Data.TaskResult.Videos[0].URL
}
return result, nil
}
+197
View File
@@ -0,0 +1,197 @@
package adapters
// * +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
// * 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 (
"encoding/json"
"errors"
"fmt"
"geekai/core/types"
"io"
"time"
"github.com/imroc/req/v3"
"gorm.io/gorm"
)
// LumaAdapter Luma 视频生成适配器
type LumaAdapter struct {
db *gorm.DB
httpClient *req.Client
}
// NewLumaAdapter 创建 Luma 适配器
func NewLumaAdapter(db *gorm.DB) *LumaAdapter {
return &LumaAdapter{
db: db,
httpClient: req.C().SetTimeout(time.Minute * 3),
}
}
// GetProvider 获取服务提供商名称
func (a *LumaAdapter) GetProvider() string {
return "luma"
}
// LumaCreateRequest Luma 创建任务请求
type LumaCreateRequest struct {
ModelName string `json:"model_name"`
UserPrompt string `json:"user_prompt"`
ExpandPrompt bool `json:"expand_prompt,omitempty"`
Loop bool `json:"loop,omitempty"`
ImageURL string `json:"image_url,omitempty"` // 图生视频
ImageEndURL string `json:"image_end_url,omitempty"` // 图生视频
Duration string `json:"duration,omitempty"` // 视频时长
Resolution string `json:"resolution,omitempty"` // 视频分辨率
}
// LumaCreateResponse Luma 创建任务响应
type LumaCreateResponse struct {
Id string `json:"id"`
Prompt string `json:"prompt"`
State string `json:"state"`
CreatedAt string `json:"created_at"`
Channel string `json:"channel,omitempty"`
}
// LumaQueryResponse Luma 查询任务响应
type LumaQueryResponse struct {
Id string `json:"id"`
State string `json:"state"`
Video struct {
URL string `json:"url"`
Width int `json:"width"`
Height int `json:"height"`
Thumbnail string `json:"thumbnail"`
DownloadURL string `json:"download_url"`
} `json:"video"`
Prompt string `json:"prompt"`
Thumbnail struct {
URL string `json:"url"`
Width int `json:"width"`
Height int `json:"height"`
} `json:"thumbnail"`
}
// CreateTask 创建视频生成任务
func (a *LumaAdapter) CreateTask(task types.VideoTask, videoConfig *types.VideoConfig) (CreateTaskResponse, error) {
// 解析任务参数
paramsMap, ok := task.Params.(map[string]any)
if !ok {
return CreateTaskResponse{}, errors.New("invalid params type for Luma video task")
}
// 构建请求参数
reqBody := LumaCreateRequest{
UserPrompt: task.Prompt,
}
// 从 params 中提取参数
if expandPrompt, ok := paramsMap["expand_prompt"].(bool); ok {
reqBody.ExpandPrompt = expandPrompt
}
if model, ok := paramsMap["model"].(string); ok {
reqBody.ModelName = model
}
if loop, ok := paramsMap["loop"].(bool); ok {
reqBody.Loop = loop
}
if imageURL, ok := paramsMap["image_url"].(string); ok {
reqBody.ImageURL = imageURL
}
if imageEndURL, ok := paramsMap["image_end_url"].(string); ok {
reqBody.ImageEndURL = imageEndURL
}
if duration, ok := paramsMap["duration"].(string); ok {
reqBody.Duration = duration
}
if resolution, ok := paramsMap["resolution"].(string); ok {
reqBody.Resolution = resolution
}
// 发送请求
apiURL := fmt.Sprintf("%s/luma/generations", videoConfig.ApiURL)
r, err := a.httpClient.R().
SetHeader("Authorization", "Bearer "+videoConfig.ApiKey).
SetBody(reqBody).
Post(apiURL)
if err != nil {
return CreateTaskResponse{}, fmt.Errorf("请求 API 出错:%v", err)
}
if r.StatusCode != 200 && r.StatusCode != 201 {
body, _ := io.ReadAll(r.Body)
return CreateTaskResponse{}, fmt.Errorf("请求 API 出错:%d, %s", r.StatusCode, string(body))
}
body, _ := io.ReadAll(r.Body)
var res LumaCreateResponse
err = json.Unmarshal(body, &res)
if err != nil {
return CreateTaskResponse{}, fmt.Errorf("解析API数据失败:%v, %s", err, string(body))
}
return CreateTaskResponse{
TaskId: res.Id,
Channel: videoConfig.ApiURL,
Prompt: res.Prompt,
State: types.VideoStatusPending,
CreatedAt: res.CreatedAt,
}, nil
}
// QueryTask 查询任务状态
func (a *LumaAdapter) QueryTask(taskId string, channel string, videoConfig *types.VideoConfig) (QueryTaskResponse, error) {
apiURL := fmt.Sprintf("%s/luma/generations/%s", videoConfig.ApiURL, taskId)
r, err := a.httpClient.R().
SetHeader("Authorization", "Bearer "+videoConfig.ApiKey).
Get(apiURL)
if err != nil {
return QueryTaskResponse{}, fmt.Errorf("请求 API 失败:%v", err)
}
defer r.Body.Close()
if r.StatusCode != 200 {
body, _ := io.ReadAll(r.Body)
return QueryTaskResponse{}, fmt.Errorf("API 返回失败:%d, %s", r.StatusCode, string(body))
}
body, _ := io.ReadAll(r.Body)
var res LumaQueryResponse
err = json.Unmarshal(body, &res)
if err != nil {
return QueryTaskResponse{}, fmt.Errorf("解析API数据失败:%v, %s", err, string(body))
}
switch res.State {
case "completed", "succeed", "success":
res.State = types.VideoStatusSuccess
case "in_progress", "running":
res.State = types.VideoStatusInProgress
case "failed":
res.State = types.VideoStatusFailed
default:
res.State = types.VideoStatusPending
}
// 构建响应
response := QueryTaskResponse{
TaskId: res.Id,
Status: res.State,
VideoURL: res.Video.DownloadURL,
Prompt: res.Prompt,
}
// 如果有原始数据,转换为 JSON 字符串
if len(body) > 0 {
response.Output = string(body)
}
return response, nil
}
@@ -0,0 +1,227 @@
package adapters
// * +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
// * 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 (
"encoding/json"
"fmt"
"geekai/core/types"
"io"
"strings"
"time"
"github.com/imroc/req/v3"
"gorm.io/gorm"
)
// MiniMaxAdapter MiniMax 视频生成适配器
type MiniMaxAdapter struct {
db *gorm.DB
httpClient *req.Client
}
// NewMiniMaxAdapter 创建 MiniMax 适配器
func NewMiniMaxAdapter(db *gorm.DB) *MiniMaxAdapter {
return &MiniMaxAdapter{
db: db,
httpClient: req.C().SetTimeout(time.Minute * 3),
}
}
// GetProvider 获取服务提供商名称
func (a *MiniMaxAdapter) GetProvider() string {
return "minimax"
}
// MiniMaxCreateRequest MiniMax 创建任务请求
type MiniMaxCreateRequest struct {
Model string `json:"model"`
Prompt string `json:"prompt"`
Duration int `json:"duration,omitempty"`
Resolution string `json:"resolution,omitempty"`
FirstFrameImage string `json:"first_frame_image,omitempty"`
LastFrameImage string `json:"last_frame_image,omitempty"`
PromptOptimizer bool `json:"prompt_optimizer,omitempty"`
}
// MiniMaxCreateResponse MiniMax 创建任务响应
type MiniMaxCreateResponse struct {
TaskId string `json:"task_id"`
BaseResp struct {
StatusCode int `json:"status_code"`
StatusMsg string `json:"status_msg"`
} `json:"base_resp"`
}
// MiniMaxFile MiniMax 文件信息
type MiniMaxFile struct {
Bytes int `json:"bytes"`
CreatedAt int64 `json:"created_at"`
DownloadURL string `json:"download_url"`
FileId int64 `json:"file_id"`
Filename string `json:"filename"`
Purpose string `json:"purpose"`
}
// MiniMaxQueryResponse MiniMax 查询任务响应
type MiniMaxQueryResponse struct {
TaskId string `json:"task_id"`
Status string `json:"status"`
FileId string `json:"file_id,omitempty"` // 顶层 file_id 可能是字符串
File *MiniMaxFile `json:"file,omitempty"` // file 对象包含详细信息
VideoWidth int `json:"video_width,omitempty"`
VideoHeight int `json:"video_height,omitempty"`
VideoURL string `json:"video_url,omitempty"`
Prompt string `json:"prompt,omitempty"`
ErrMsg string `json:"err_msg,omitempty"`
StatusMsg string `json:"status_msg,omitempty"`
BaseResp struct {
StatusCode int `json:"status_code"`
StatusMsg string `json:"status_msg"`
} `json:"base_resp"`
}
// CreateTask 创建视频生成任务
func (a *MiniMaxAdapter) CreateTask(task types.VideoTask, videoConfig *types.VideoConfig) (CreateTaskResponse, error) {
// 解析任务参数
paramsMap, ok := task.Params.(map[string]interface{})
if !ok {
return CreateTaskResponse{}, fmt.Errorf("invalid params type for MiniMax video task")
}
// 构建请求参数
reqBody := MiniMaxCreateRequest{
Prompt: task.Prompt,
}
// 从 params 中提取参数
if model, ok := paramsMap["model"].(string); ok {
reqBody.Model = model
}
if duration, ok := paramsMap["duration"].(float64); ok {
reqBody.Duration = int(duration)
} else if duration, ok := paramsMap["duration"].(int); ok {
reqBody.Duration = duration
}
if resolution, ok := paramsMap["resolution"].(string); ok {
reqBody.Resolution = resolution
}
if firstFrameImage, ok := paramsMap["first_frame_image"].(string); ok {
reqBody.FirstFrameImage = firstFrameImage
}
if lastFrameImage, ok := paramsMap["last_frame_image"].(string); ok {
reqBody.LastFrameImage = lastFrameImage
}
if promptOptimizer, ok := paramsMap["prompt_optimizer"].(bool); ok {
reqBody.PromptOptimizer = promptOptimizer
} else {
reqBody.PromptOptimizer = true // 默认值
}
// 发送请求
apiURL := fmt.Sprintf("%s/minimax/v1/video_generation", videoConfig.ApiURL)
r, err := a.httpClient.R().
SetHeader("Authorization", "Bearer "+videoConfig.ApiKey).
SetHeader("Content-Type", "application/json").
SetBody(reqBody).
Post(apiURL)
if err != nil {
return CreateTaskResponse{}, fmt.Errorf("请求 API 出错:%v", err)
}
if r.StatusCode != 200 && r.StatusCode != 201 {
body, _ := io.ReadAll(r.Body)
return CreateTaskResponse{}, fmt.Errorf("请求 API 出错:%d, %s", r.StatusCode, string(body))
}
body, _ := io.ReadAll(r.Body)
var res MiniMaxCreateResponse
err = json.Unmarshal(body, &res)
if err != nil {
return CreateTaskResponse{}, fmt.Errorf("解析API数据失败:%v, %s", err, string(body))
}
if res.BaseResp.StatusCode != 0 {
return CreateTaskResponse{}, fmt.Errorf("API 返回错误:%s", res.BaseResp.StatusMsg)
}
return CreateTaskResponse{
TaskId: res.TaskId,
Channel: videoConfig.ApiURL,
Prompt: task.Prompt,
State: types.VideoStatusPending,
CreatedAt: time.Now().Format(time.RFC3339),
}, nil
}
// QueryTask 查询任务状态
func (a *MiniMaxAdapter) QueryTask(taskId string, channel string, videoConfig *types.VideoConfig) (QueryTaskResponse, error) {
// MiniMax 查询接口
apiURL := fmt.Sprintf("%s/minimax/v1/query/video_generation?task_id=%s", videoConfig.ApiURL, taskId)
r, err := a.httpClient.R().
SetHeader("Authorization", "Bearer "+videoConfig.ApiKey).
Get(apiURL)
if err != nil {
return QueryTaskResponse{}, fmt.Errorf("请求 API 失败:%v", err)
}
defer r.Body.Close()
if r.StatusCode != 200 {
body, _ := io.ReadAll(r.Body)
return QueryTaskResponse{}, fmt.Errorf("API 返回失败:%d, %s", r.StatusCode, string(body))
}
body, _ := io.ReadAll(r.Body)
var res MiniMaxQueryResponse
err = json.Unmarshal(body, &res)
if err != nil {
return QueryTaskResponse{}, fmt.Errorf("解析API数据失败:%v, %s", err, string(body))
}
if res.BaseResp.StatusCode != 0 {
return QueryTaskResponse{}, fmt.Errorf("API 返回错误:%s", res.BaseResp.StatusMsg)
}
// 转换状态(处理大小写)
state := strings.ToLower(res.Status)
switch state {
case "completed", "succeed", "success":
state = types.VideoStatusSuccess
case "in_progress", "running":
state = types.VideoStatusInProgress
case "failed":
state = types.VideoStatusFailed
default:
state = types.VideoStatusPending
}
// 获取视频URL,优先从 file.download_url 获取
videoURL := res.VideoURL
if res.File != nil && res.File.DownloadURL != "" {
videoURL = res.File.DownloadURL
}
// 构建响应
response := QueryTaskResponse{
TaskId: res.TaskId,
Status: state,
VideoURL: videoURL,
Prompt: res.Prompt,
ErrMsg: res.ErrMsg,
StatusMsg: res.StatusMsg,
}
// 如果有原始数据,转换为 JSON 字符串
if len(body) > 0 {
response.Output = string(body)
}
return response, nil
}
+347
View File
@@ -0,0 +1,347 @@
package adapters
// * +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
// * 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 (
"bytes"
"context"
"encoding/json"
"fmt"
"geekai/core/types"
"io"
"mime/multipart"
"net/http"
"geekai/utils"
"time"
"github.com/imroc/req/v3"
"gorm.io/gorm"
)
// SoraAdapter Sora 视频生成适配器
type SoraAdapter struct {
db *gorm.DB
httpClient *req.Client
}
// NewSoraAdapter 创建 Sora 适配器
func NewSoraAdapter(db *gorm.DB) *SoraAdapter {
return &SoraAdapter{
db: db,
httpClient: req.C().SetTimeout(time.Minute * 3),
}
}
// GetProvider 获取服务提供商名称
func (a *SoraAdapter) GetProvider() string {
return "sora"
}
// SoraCreateRequest Sora 创建任务请求
type SoraCreateRequest struct {
Model string `json:"model"` // 模型名称:sora-2, sora-2-pro
Prompt string `json:"prompt"` // 提示词
Size string `json:"size,omitempty"` // 分辨率:1280x720, 720x1280, 1792x1024, 1024x1792
InputReference interface{} `json:"input_reference,omitempty"` // 图生视频的参考图片,官方为对象 {"image_url": "..."},也兼容字符串 URL
Seconds string `json:"seconds,omitempty"` // 视频时长(秒),默认4秒
Watermark bool `json:"watermark,omitempty"` // 是否添加水印
}
// SoraCreateResponse Sora 创建任务响应
type SoraCreateResponse struct {
ID string `json:"id"` // 任务ID
Object string `json:"object"` // 对象类型,固定为 "video"
Model string `json:"model"` // 模型名称
Status string `json:"status"` // 状态:queued, in_progress, completed, failed
CreatedAt int64 `json:"created_at"` // 创建时间戳
Seconds string `json:"seconds"` // 视频时长
Size string `json:"size"` // 分辨率
Error *SoraError `json:"error,omitempty"` // 错误信息(成功时为null
}
// SoraQueryResponse Sora 查询任务响应
type SoraQueryResponse struct {
ID string `json:"id"` // 任务ID
Object string `json:"object"` // 对象类型,固定为 "video"
Model string `json:"model"` // 模型名称
Status string `json:"status"` // 状态:queued, in_progress, completed, failed
Progress int `json:"progress"` // 进度(0-100
CreatedAt int64 `json:"created_at"` // 创建时间戳
Seconds string `json:"seconds"` // 视频时长
Size string `json:"size"` // 分辨率
Error *SoraError `json:"error,omitempty"` // 错误信息(成功时为null
VideoURL string `json:"video_url"` // 视频URL(成功时生成)
}
type SoraError struct {
Code string `json:"code"`
Message string `json:"message"`
}
// CreateTask 创建视频生成任务
func (a *SoraAdapter) CreateTask(task types.VideoTask, videoConfig *types.VideoConfig) (CreateTaskResponse, error) {
// 解析任务参数
paramsMap, ok := task.Params.(map[string]any)
if !ok {
return CreateTaskResponse{}, fmt.Errorf("invalid params type for Sora video task")
}
// 是否调用官方 Sora 接口
isOfficial := false
if v, ok := paramsMap["is_official"].(bool); ok {
isOfficial = v
}
// 提取通用参数
model, ok := paramsMap["model"].(string)
if !ok || model == "" {
return CreateTaskResponse{}, fmt.Errorf("model 参数必填")
}
size, _ := paramsMap["size"].(string)
seconds := "10" // 默认 10 秒
if v, ok := paramsMap["seconds"].(string); ok && v != "" {
seconds = v
} else if duration, ok := paramsMap["duration"].(float64); ok {
seconds = fmt.Sprintf("%.0f", duration)
} else if duration, ok := paramsMap["duration"].(int); ok {
seconds = fmt.Sprintf("%d", duration)
}
watermark := false
if v, ok := paramsMap["watermark"].(bool); ok {
watermark = v
}
// 处理图生视频(input_reference 参数)
// 支持单个字符串或数组的第一个元素
var imageURL string
if inputRef, ok := paramsMap["input_reference"].(string); ok && inputRef != "" {
imageURL = inputRef
} else if images, ok := paramsMap["images"].([]interface{}); ok && len(images) > 0 {
// 兼容旧的 images 参数格式
if imgStr, ok := images[0].(string); ok && imgStr != "" {
imageURL = imgStr
}
} else if image, ok := paramsMap["image"].(string); ok && image != "" {
// 兼容 image 参数
imageURL = image
}
// 官方 Sora:使用 multipart/form-data 携带文件
if isOfficial && imageURL != "" {
return a.createOfficialSoraTask(task, videoConfig, model, size, seconds, imageURL)
}
// 其他场景:保持原来的 JSON 调用,input_reference 继续传 URL 字符串
reqBody := SoraCreateRequest{
Model: model,
Prompt: task.Prompt,
Size: size,
Seconds: seconds,
Watermark: watermark,
}
if imageURL != "" {
reqBody.InputReference = imageURL
}
// 发送 JSON 请求
apiURL := fmt.Sprintf("%s/v1/videos", videoConfig.ApiURL)
r, err := a.httpClient.R().
SetHeader("Authorization", "Bearer "+videoConfig.ApiKey).
SetHeader("Content-Type", "application/json").
SetBody(reqBody).
Post(apiURL)
if err != nil {
return CreateTaskResponse{}, fmt.Errorf("请求 API 出错:%v", err)
}
if r.StatusCode != 200 && r.StatusCode != 201 {
body, _ := io.ReadAll(r.Body)
return CreateTaskResponse{}, fmt.Errorf("请求 API 出错:%d, %s", r.StatusCode, string(body))
}
body, _ := io.ReadAll(r.Body)
var res SoraCreateResponse
err = json.Unmarshal(body, &res)
if err != nil {
return CreateTaskResponse{}, fmt.Errorf("解析API数据失败:%v, %s", err, string(body))
}
// 转换状态:queued -> pending
state := res.Status
if state == "queued" || state == "in_progress" || state == "" {
state = "pending"
}
return CreateTaskResponse{
TaskId: res.ID,
Channel: videoConfig.ApiURL,
Prompt: task.Prompt,
State: state,
CreatedAt: time.Unix(res.CreatedAt, 0).Format(time.RFC3339),
}, nil
}
// createOfficialSoraTask 调用官方 Sora API,使用 multipart/form-data 携带图片文件
func (a *SoraAdapter) createOfficialSoraTask(task types.VideoTask, videoConfig *types.VideoConfig, model, size, seconds, imageURL string) (CreateTaskResponse, error) {
if videoConfig == nil || videoConfig.ApiURL == "" || videoConfig.ApiKey == "" {
return CreateTaskResponse{}, fmt.Errorf("Sora 视频配置不完整")
}
imgData, err := downloadImageBytes(imageURL)
if err != nil {
return CreateTaskResponse{}, fmt.Errorf("下载参考图片失败:%v", err)
}
var buf bytes.Buffer
writer := multipart.NewWriter(&buf)
// 文本字段
if err = writer.WriteField("prompt", task.Prompt); err != nil {
return CreateTaskResponse{}, err
}
if err = writer.WriteField("model", model); err != nil {
return CreateTaskResponse{}, err
}
if size != "" {
if err = writer.WriteField("size", size); err != nil {
return CreateTaskResponse{}, err
}
}
if seconds != "" {
if err = writer.WriteField("seconds", seconds); err != nil {
return CreateTaskResponse{}, err
}
}
// 文件字段
fileWriter, err := writer.CreateFormFile("input_reference", "image")
if err != nil {
return CreateTaskResponse{}, err
}
if _, err = fileWriter.Write(imgData); err != nil {
return CreateTaskResponse{}, err
}
if err = writer.Close(); err != nil {
return CreateTaskResponse{}, err
}
apiURL := fmt.Sprintf("%s/v1/videos", videoConfig.ApiURL)
req, err := http.NewRequest(http.MethodPost, apiURL, &buf)
if err != nil {
return CreateTaskResponse{}, err
}
req.Header.Set("Authorization", "Bearer "+videoConfig.ApiKey)
req.Header.Set("Content-Type", writer.FormDataContentType())
client := &http.Client{Timeout: 3 * time.Minute}
resp, err := client.Do(req)
if err != nil {
return CreateTaskResponse{}, fmt.Errorf("请求官方 Sora API 出错:%v", err)
}
defer resp.Body.Close()
body, _ := io.ReadAll(resp.Body)
if resp.StatusCode != http.StatusOK && resp.StatusCode != http.StatusCreated {
return CreateTaskResponse{}, fmt.Errorf("请求官方 Sora API 出错:%d, %s", resp.StatusCode, string(body))
}
var res SoraCreateResponse
if err = json.Unmarshal(body, &res); err != nil {
return CreateTaskResponse{}, fmt.Errorf("解析官方 Sora API 数据失败:%v, %s", err, string(body))
}
state := res.Status
if state == "queued" || state == "in_progress" || state == "" {
state = "pending"
}
return CreateTaskResponse{
TaskId: res.ID,
Channel: videoConfig.ApiURL,
Prompt: task.Prompt,
State: state,
CreatedAt: time.Unix(res.CreatedAt, 0).Format(time.RFC3339),
}, nil
}
// downloadImageBytes 下载远程图片并返回二进制内容,用于 multipart 文件上传
func downloadImageBytes(imageURL string) ([]byte, error) {
body, _, err := utils.FetchURLBytes(context.Background(), imageURL, "", 3*time.Minute, 2, 32<<20)
return body, err
}
// downloadImageAsDataURL 下载远程图片并转为 data URL,避免向官方 Sora 直接传地址
// QueryTask 查询任务状态
func (a *SoraAdapter) QueryTask(taskId string, channel string, videoConfig *types.VideoConfig) (QueryTaskResponse, error) {
apiURL := fmt.Sprintf("%s/v1/videos/%s", videoConfig.ApiURL, taskId)
r, err := a.httpClient.R().
SetHeader("Authorization", "Bearer "+videoConfig.ApiKey).
Get(apiURL)
if err != nil {
return QueryTaskResponse{}, fmt.Errorf("请求 API 失败:%v", err)
}
defer r.Body.Close()
if r.StatusCode != 200 {
body, _ := io.ReadAll(r.Body)
return QueryTaskResponse{}, fmt.Errorf("API 返回失败:%d, %s", r.StatusCode, string(body))
}
body, _ := io.ReadAll(r.Body)
var res SoraQueryResponse
err = json.Unmarshal(body, &res)
if err != nil {
return QueryTaskResponse{}, fmt.Errorf("解析API数据失败:%v, %s", err, string(body))
}
// 转换状态:queued -> pending, completed -> success
state := res.Status
switch state {
case "completed", "succeed", "success":
state = types.VideoStatusSuccess
case "in_progress", "running":
state = types.VideoStatusInProgress
case "failed":
state = types.VideoStatusFailed
default:
state = types.VideoStatusPending
}
// 处理错误信息
errMsg := ""
if res.Error != nil {
errMsg = res.Error.Message
} else {
errMsg = fmt.Sprintf("进度: %d%%", res.Progress)
}
// 构建响应
response := QueryTaskResponse{
TaskId: res.ID,
Status: state,
Progress: res.Progress,
VideoURL: res.VideoURL,
Prompt: "", // Sora API 响应中不包含 prompt 字段
ErrMsg: errMsg,
StatusMsg: fmt.Sprintf("进度: %d%%", res.Progress),
}
// 如果有原始数据,转换为 JSON 字符串
if len(body) > 0 {
response.Output = string(body)
}
return response, nil
}
+226
View File
@@ -0,0 +1,226 @@
package adapters
// * +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
// * 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 (
"encoding/json"
"fmt"
"geekai/core/types"
"io"
"strings"
"time"
"github.com/imroc/req/v3"
"gorm.io/gorm"
)
// VideoAdapter 视频生成适配器接口
type VideoAdapter interface {
// CreateTask 创建视频生成任务
CreateTask(task types.VideoTask, videoConfig *types.VideoConfig) (CreateTaskResponse, error)
// QueryTask 查询任务状态
QueryTask(taskId string, channel string, videoConfig *types.VideoConfig) (QueryTaskResponse, error)
// GetProvider 获取服务提供商名称(不带版本号:veo, sora, luma
GetProvider() string
}
// VeoAdapter Veo 视频生成适配器
type VeoAdapter struct {
db *gorm.DB
httpClient *req.Client
}
// NewVeoAdapter 创建 Veo 适配器
func NewVeoAdapter(db *gorm.DB) *VeoAdapter {
return &VeoAdapter{
db: db,
httpClient: req.C().SetTimeout(time.Minute * 3),
}
}
// GetProvider 获取服务提供商名称
func (a *VeoAdapter) GetProvider() string {
return "veo"
}
// VeoCreateRequest Veo 创建任务请求
type VeoCreateRequest struct {
Prompt string `json:"prompt"`
Model string `json:"model"`
EnhancePrompt bool `json:"enhance_prompt,omitempty"`
EnableUpsample bool `json:"enable_upsample,omitempty"`
AspectRatio string `json:"aspect_ratio,omitempty"`
Images []string `json:"images,omitempty"` // 图生视频时使用
}
// VeoCreateResponse Veo 创建任务响应
type VeoCreateResponse struct {
TaskId string `json:"task_id"`
}
// VeoQueryResponse Veo 查询任务响应
type VeoQueryResponse struct {
TaskId string `json:"task_id"`
Platform string `json:"platform"`
Action string `json:"action"`
Status string `json:"status"`
FailReason string `json:"fail_reason"`
SubmitTime int64 `json:"submit_time"`
StartTime int64 `json:"start_time"`
FinishTime int64 `json:"finish_time"`
Progress string `json:"progress"`
Data VeoQueryData `json:"data"`
SearchItem string `json:"search_item"`
}
// VeoQueryData Veo 查询响应中的 data 字段
type VeoQueryData struct {
Output string `json:"output"`
}
// CreateTask 创建视频生成任务
func (a *VeoAdapter) CreateTask(task types.VideoTask, videoConfig *types.VideoConfig) (CreateTaskResponse, error) {
// 解析任务参数
paramsMap, ok := task.Params.(map[string]any)
if !ok {
return CreateTaskResponse{}, fmt.Errorf("invalid params type for Veo video task")
}
// 构建请求参数
reqBody := VeoCreateRequest{
Prompt: task.Prompt,
}
// 从 params 中提取参数
if model, ok := paramsMap["model"].(string); ok {
reqBody.Model = model
}
if enhancePrompt, ok := paramsMap["enhance_prompt"].(bool); ok {
reqBody.EnhancePrompt = enhancePrompt
}
if enableUpsample, ok := paramsMap["enable_upsample"].(bool); ok {
reqBody.EnableUpsample = enableUpsample
}
if aspectRatio, ok := paramsMap["aspect_ratio"].(string); ok {
reqBody.AspectRatio = aspectRatio
}
// 处理图生视频(images 参数)
if images, ok := paramsMap["images"].([]interface{}); ok {
imageUrls := make([]string, 0)
for _, img := range images {
if imgStr, ok := img.(string); ok {
imageUrls = append(imageUrls, imgStr)
}
}
if len(imageUrls) > 0 {
reqBody.Images = imageUrls
}
}
// 发送请求
apiURL := fmt.Sprintf("%s/v2/videos/generations", videoConfig.ApiURL)
r, err := a.httpClient.R().
SetHeader("Authorization", "Bearer "+videoConfig.ApiKey).
SetHeader("Content-Type", "application/json").
SetBody(reqBody).
Post(apiURL)
if err != nil {
return CreateTaskResponse{}, fmt.Errorf("请求 API 出错:%v", err)
}
if r.StatusCode != 200 && r.StatusCode != 201 {
body, _ := io.ReadAll(r.Body)
return CreateTaskResponse{}, fmt.Errorf("请求 API 出错:%d, %s", r.StatusCode, string(body))
}
body, _ := io.ReadAll(r.Body)
var res VeoCreateResponse
err = json.Unmarshal(body, &res)
if err != nil {
return CreateTaskResponse{}, fmt.Errorf("解析API数据失败:%v, %s", err, string(body))
}
return CreateTaskResponse{
TaskId: res.TaskId,
Channel: videoConfig.ApiURL,
Prompt: task.Prompt,
State: types.VideoStatusPending,
CreatedAt: time.Now().Format(time.RFC3339),
}, nil
}
// QueryTask 查询任务状态
func (a *VeoAdapter) QueryTask(taskId string, channel string, videoConfig *types.VideoConfig) (QueryTaskResponse, error) {
apiURL := fmt.Sprintf("%s/v2/videos/generations/%s", videoConfig.ApiURL, taskId)
r, err := a.httpClient.R().
SetHeader("Authorization", "Bearer "+videoConfig.ApiKey).
Get(apiURL)
if err != nil {
return QueryTaskResponse{}, fmt.Errorf("请求 API 失败:%v", err)
}
defer r.Body.Close()
if r.StatusCode != 200 {
body, _ := io.ReadAll(r.Body)
return QueryTaskResponse{}, fmt.Errorf("API 返回失败:%d, %s", r.StatusCode, string(body))
}
body, _ := io.ReadAll(r.Body)
var res VeoQueryResponse
err = json.Unmarshal(body, &res)
if err != nil {
return QueryTaskResponse{}, fmt.Errorf("解析API数据失败:%v, %s", err, string(body))
}
// 转换状态(SUCCESS -> success, FAILED -> failed, 其他保持原样)
state := strings.ToLower(res.Status)
switch state {
case "in_progress", "running":
state = types.VideoStatusInProgress
case "completed", "succeed", "success":
state = types.VideoStatusSuccess
case "failed":
state = types.VideoStatusFailed
default:
state = types.VideoStatusPending
}
// 解析进度(从 "100%" 转换为 100
progress := 0
if res.Progress != "" {
// 移除 % 符号并转换为整数
progressStr := strings.TrimSuffix(res.Progress, "%")
if p, err := fmt.Sscanf(progressStr, "%d", &progress); err == nil && p == 1 {
// 成功解析
}
}
// 从 data.output 中提取视频 URL
videoURL := res.Data.Output
// 构建响应
response := QueryTaskResponse{
TaskId: res.TaskId,
Status: state,
Progress: progress,
VideoURL: videoURL,
ErrMsg: res.FailReason,
}
// 如果有原始数据,转换为 JSON 字符串
if len(body) > 0 {
response.Output = string(body)
}
return response, nil
}
+220
View File
@@ -0,0 +1,220 @@
package adapters
// * +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
// * 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 (
"encoding/json"
"fmt"
"geekai/core/types"
"io"
"strings"
"time"
"github.com/imroc/req/v3"
"gorm.io/gorm"
)
// WanAdapter Wan(通义万相)视频生成适配器
type WanAdapter struct {
db *gorm.DB
httpClient *req.Client
}
// NewWanAdapter 创建 Wan 适配器
func NewWanAdapter(db *gorm.DB) *WanAdapter {
return &WanAdapter{
db: db,
httpClient: req.C().SetTimeout(time.Minute * 3),
}
}
// GetProvider 获取服务提供商名称
func (a *WanAdapter) GetProvider() string {
return "wan"
}
// WanCreateRequest Wan 创建任务请求
type WanCreateRequest struct {
Prompt string `json:"prompt"`
Model string `json:"model"`
Duration int `json:"duration,omitempty"`
Resolution string `json:"resolution,omitempty"`
NegativePrompt string `json:"negative_prompt,omitempty"`
Images []string `json:"images,omitempty"`
PromptExtend bool `json:"prompt_extend,omitempty"`
}
// WanCreateResponse Wan 创建任务响应
type WanCreateResponse struct {
TaskId string `json:"task_id"`
}
// WanQueryResponse Wan 查询任务响应
type WanQueryResponse struct {
TaskId string `json:"task_id"`
Platform string `json:"platform"`
Action string `json:"action"`
Status string `json:"status"`
FailReason string `json:"fail_reason"`
SubmitTime int64 `json:"submit_time"`
StartTime int64 `json:"start_time"`
FinishTime int64 `json:"finish_time"`
Progress string `json:"progress"`
Data WanQueryData `json:"data"`
SearchItem string `json:"search_item"`
}
// WanQueryData Wan 查询响应中的 data 字段
type WanQueryData struct {
Output string `json:"output"`
}
// CreateTask 创建视频生成任务
func (a *WanAdapter) CreateTask(task types.VideoTask, videoConfig *types.VideoConfig) (CreateTaskResponse, error) {
// 解析任务参数
paramsMap, ok := task.Params.(map[string]any)
if !ok {
return CreateTaskResponse{}, fmt.Errorf("invalid params type for Wan video task")
}
// 构建请求参数
reqBody := WanCreateRequest{
Prompt: task.Prompt,
}
// 从 params 中提取参数
if model, ok := paramsMap["model"].(string); ok {
reqBody.Model = model
}
if duration, ok := paramsMap["duration"].(float64); ok {
reqBody.Duration = int(duration)
} else if duration, ok := paramsMap["duration"].(int); ok {
reqBody.Duration = duration
}
if resolution, ok := paramsMap["resolution"].(string); ok {
reqBody.Resolution = resolution
}
if images, ok := paramsMap["images"].([]any); ok {
imageUrls := make([]string, 0)
for _, img := range images {
if imgStr, ok := img.(string); ok {
imageUrls = append(imageUrls, imgStr)
}
}
if len(imageUrls) > 0 {
reqBody.Images = imageUrls
}
}
if negativePrompt, ok := paramsMap["negative_prompt"].(string); ok {
reqBody.NegativePrompt = negativePrompt
}
if promptExtend, ok := paramsMap["prompt_extend"].(bool); ok {
reqBody.PromptExtend = promptExtend
}
logger.Debugf("WanCreateRequest: %+v", reqBody)
// 发送请求
apiURL := fmt.Sprintf("%s/v2/videos/generations", videoConfig.ApiURL)
r, err := a.httpClient.R().
SetHeader("Authorization", "Bearer "+videoConfig.ApiKey).
SetHeader("Content-Type", "application/json").
SetBody(reqBody).
Post(apiURL)
if err != nil {
return CreateTaskResponse{}, fmt.Errorf("请求 API 出错:%v", err)
}
if r.StatusCode != 200 && r.StatusCode != 201 {
body, _ := io.ReadAll(r.Body)
return CreateTaskResponse{}, fmt.Errorf("请求 API 出错:%d, %s", r.StatusCode, string(body))
}
body, _ := io.ReadAll(r.Body)
var res WanCreateResponse
err = json.Unmarshal(body, &res)
if err != nil {
return CreateTaskResponse{}, fmt.Errorf("解析API数据失败:%v, %s", err, string(body))
}
return CreateTaskResponse{
TaskId: res.TaskId,
Channel: videoConfig.ApiURL,
Prompt: task.Prompt,
State: "pending",
CreatedAt: time.Now().Format(time.RFC3339),
}, nil
}
// QueryTask 查询任务状态
func (a *WanAdapter) QueryTask(taskId string, channel string, videoConfig *types.VideoConfig) (QueryTaskResponse, error) {
apiURL := fmt.Sprintf("%s/v2/videos/generations/%s", videoConfig.ApiURL, taskId)
r, err := a.httpClient.R().
SetHeader("Authorization", "Bearer "+videoConfig.ApiKey).
Get(apiURL)
if err != nil {
return QueryTaskResponse{}, fmt.Errorf("请求 API 失败:%v", err)
}
defer r.Body.Close()
if r.StatusCode != 200 {
body, _ := io.ReadAll(r.Body)
return QueryTaskResponse{}, fmt.Errorf("API 返回失败:%d, %s", r.StatusCode, string(body))
}
body, _ := io.ReadAll(r.Body)
var res WanQueryResponse
err = json.Unmarshal(body, &res)
if err != nil {
return QueryTaskResponse{}, fmt.Errorf("解析API数据失败:%v, %s", err, string(body))
}
// 转换状态(SUCCESS -> success, FAILED -> failed, 其他保持原样)
state := strings.ToLower(res.Status)
switch state {
case "in_progress", "running":
state = types.VideoStatusInProgress
case "completed", "succeed", "success":
state = types.VideoStatusSuccess
case "failed", "failure":
state = types.VideoStatusFailed
default:
state = types.VideoStatusPending
}
// 解析进度(从 "100%" 转换为 100
progress := 0
if res.Progress != "" {
// 移除 % 符号并转换为整数
progressStr := strings.TrimSuffix(res.Progress, "%")
if p, err := fmt.Sscanf(progressStr, "%d", &progress); err == nil && p == 1 {
// 成功解析
}
}
// 从 data.output 中提取视频 URL
videoURL := res.Data.Output
// 构建响应
response := QueryTaskResponse{
TaskId: res.TaskId,
Status: state,
Progress: progress,
VideoURL: videoURL,
ErrMsg: res.FailReason,
}
// 如果有原始数据,转换为 JSON 字符串
if len(body) > 0 {
response.Output = string(body)
}
return response, nil
}
+77
View File
@@ -0,0 +1,77 @@
package video
// * +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
// * 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 (
"errors"
"fmt"
"geekai/core/types"
"geekai/store/model"
"geekai/utils"
"gorm.io/gorm"
)
// GetVideoConfig 从数据库获取视频配置
func GetVideoConfig(db *gorm.DB) (*types.VideoConfig, error) {
var config model.Config
err := db.Where("name", types.ConfigKeyVideo).First(&config).Error
if err != nil {
if err == gorm.ErrRecordNotFound {
return nil, errors.New("视频配置不存在,请在管理后台配置")
}
return nil, fmt.Errorf("获取视频配置失败: %v", err)
}
var videoConfig types.VideoConfig
err = utils.JsonDecode(config.Value, &videoConfig)
if err != nil {
return nil, fmt.Errorf("解析视频配置失败: %v", err)
}
return &videoConfig, nil
}
// GetModelPowerConfig 获取指定模型的算力配置
func GetModelPowerConfig(db *gorm.DB, modelKey string) (*types.VideoModelPower, error) {
config, err := GetVideoConfig(db)
if err != nil {
return nil, err
}
modelPower, ok := config.VideoPowers[modelKey]
if !ok {
return nil, fmt.Errorf("模型 %s 的算力配置不存在", modelKey)
}
return &modelPower, nil
}
// CalculatePower 根据 modelKey 和 priceKey 计算算力
// modelKey: 模型标识(如 "veo-2.0", "sora-2.0"
// priceKey: 价格键(如 "fixed", "5_720P", "std_5_sound" 等)
func CalculatePower(db *gorm.DB, modelKey string, priceKey string) (int, error) {
if priceKey == "" {
return 0, errors.New("priceKey 不能为空")
}
modelPower, err := GetModelPowerConfig(db, modelKey)
if err != nil {
return 0, err
}
power, ok := modelPower.PowerConfig[priceKey]
if !ok {
return 0, fmt.Errorf("模型 %s 的价格配置 %s 不存在", modelKey, priceKey)
}
if power <= 0 {
return 0, fmt.Errorf("模型 %s 的价格配置 %s 的值无效", modelKey, priceKey)
}
return power, nil
}
-663
View File
@@ -1,663 +0,0 @@
package video
// * +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
// * 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 (
"bytes"
"encoding/json"
"errors"
"fmt"
"geekai/core/types"
logger2 "geekai/logger"
"geekai/service"
"geekai/service/oss"
"geekai/store"
"geekai/store/model"
"geekai/utils"
"io"
"net/http"
"time"
"github.com/go-redis/redis/v8"
"github.com/imroc/req/v3"
"gorm.io/gorm"
)
var logger = logger2.GetLogger()
type Service struct {
httpClient *req.Client
db *gorm.DB
uploadManager *oss.UploaderManager
taskQueue *store.RedisQueue
userService *service.UserService
}
func NewService(db *gorm.DB, manager *oss.UploaderManager, redisCli *redis.Client, userService *service.UserService) *Service {
return &Service{
httpClient: req.C().SetTimeout(time.Minute * 3),
db: db,
taskQueue: store.NewRedisQueue("Video_Task_Queue", redisCli),
uploadManager: manager,
userService: userService,
}
}
func (s *Service) PushTask(task types.VideoTask) {
logger.Infof("add a new Video task to the task list: %+v", task)
if err := s.taskQueue.RPush(task); err != nil {
logger.Errorf("push video task to queue failed: %v", err)
}
}
func (s *Service) Run() {
// 将数据库中未提交的任务加载到队列
var jobs []model.VideoJob
s.db.Where("task_id", "").Where("progress", 0).Find(&jobs)
for _, v := range jobs {
var task types.VideoTask
err := utils.JsonDecode(v.TaskInfo, &task)
if err != nil {
logger.Errorf("decode task info with error: %v", err)
continue
}
task.Id = v.Id
s.PushTask(task)
}
logger.Info("Starting Video job consumer...")
go func() {
for {
var task types.VideoTask
err := s.taskQueue.LPop(&task)
if err != nil {
logger.Errorf("taking task with error: %v", err)
continue
}
if task.Type == types.VideoLuma {
// translate prompt
if utils.HasChinese(task.Prompt) {
content, err := utils.OpenAIRequest(s.db, fmt.Sprintf(service.TranslatePromptTemplate, task.Prompt), task.TranslateModelId)
if err == nil {
task.Prompt = content
} else {
logger.Warnf("error with translate prompt: %v", err)
}
}
var r LumaRespVo
r, err = s.LumaCreate(task)
if err != nil {
logger.Errorf("create task with error: %v", err)
err = s.db.Model(&model.VideoJob{Id: task.Id}).UpdateColumns(map[string]interface{}{
"err_msg": err.Error(),
"progress": service.FailTaskProgress,
"cover_url": "/images/failed.jpg",
}).Error
if err != nil {
logger.Errorf("update task with error: %v", err)
}
continue
}
// 更新任务信息
err = s.db.Model(&model.VideoJob{Id: task.Id}).UpdateColumns(map[string]interface{}{
"task_id": r.Id,
"channel": r.Channel,
"prompt_ext": r.Prompt,
}).Error
if err != nil {
logger.Errorf("update task with error: %v", err)
s.PushTask(task)
}
} else if task.Type == types.VideoKeLing {
var r KeLingRespVo
r, err = s.KeLingCreate(task)
logger.Debugf("ke ling create task result: %+v", r)
if err != nil {
logger.Errorf("create task with error: %v", err)
err = s.db.Model(&model.VideoJob{Id: task.Id}).UpdateColumns(map[string]interface{}{
"err_msg": err.Error(),
"progress": service.FailTaskProgress,
"cover_url": "/images/failed.jpg",
}).Error
if err != nil {
logger.Errorf("update task with error: %v", err)
}
continue
}
// 更新任务信息
err = s.db.Model(&model.VideoJob{Id: task.Id}).UpdateColumns(map[string]interface{}{
"task_id": r.Data.TaskID,
"channel": r.Channel,
"prompt_ext": task.Prompt,
}).Error
if err != nil {
logger.Errorf("update task with error: %v", err)
s.PushTask(task)
}
}
}
}()
}
func (s *Service) DownloadFiles() {
go func() {
var items []model.VideoJob
for {
res := s.db.Where("progress", 102).Find(&items)
if res.Error != nil {
continue
}
for _, v := range items {
if v.WaterURL == "" {
continue
}
logger.Infof("try download video: %s", v.WaterURL)
videoURL, err := s.uploadManager.GetUploadHandler().PutUrlFile(v.WaterURL, ".mp4", true)
if err != nil {
logger.Errorf("download video with error: %v", err)
continue
}
logger.Infof("download video success: %s", videoURL)
v.WaterURL = videoURL
if v.VideoURL != "" {
logger.Infof("try download no water video: %s", v.VideoURL)
videoURL, err = s.uploadManager.GetUploadHandler().PutUrlFile(v.VideoURL, ".mp4", true)
if err != nil {
logger.Errorf("download video with error: %v", err)
continue
}
}
logger.Infof("download no water video success: %s", videoURL)
v.VideoURL = videoURL
v.Progress = 100
s.db.Updates(&v)
// Convert TaskInfo to VideoTask
var videoTask types.VideoTask
if err := json.Unmarshal([]byte(v.TaskInfo), &videoTask); err != nil {
logger.Errorf("failed to unmarshal task info to VideoTask: %v", err)
continue
}
}
time.Sleep(time.Second * 10)
}
}()
}
// SyncTaskProgress 异步拉取任务
func (s *Service) SyncTaskProgress() {
go func() {
var jobs []model.VideoJob
for {
res := s.db.Where("progress < ?", 100).Where("task_id <> ?", "").Find(&jobs)
if res.Error != nil {
continue
}
for _, job := range jobs {
if job.Type == types.VideoLuma {
task, err := s.QueryLumaTask(job.TaskId, job.Channel)
if err != nil {
logger.Errorf("query task with error: %v", err)
// 更新任务信息
s.db.Model(&model.VideoJob{Id: job.Id}).UpdateColumns(map[string]interface{}{
"progress": service.FailTaskProgress, // 102 表示资源未下载完成,
"err_msg": err.Error(),
"cover_url": "/images/failed.jpg",
})
continue
}
logger.Debugf("task: %+v", task)
if task.State == "completed" { // 更新任务信息
data := map[string]interface{}{
"progress": 102, // 102 表示资源未下载完成,
"water_url": task.Video.Url,
"raw_data": utils.JsonEncode(task),
"prompt_ext": task.Prompt,
"cover_url": task.Thumbnail.Url,
}
if task.Video.DownloadUrl != "" {
data["video_url"] = task.Video.DownloadUrl
}
err = s.db.Model(&model.VideoJob{Id: job.Id}).UpdateColumns(data).Error
if err != nil {
logger.Errorf("更新数据库失败:%v", err)
continue
}
}
} else if job.Type == types.VideoKeLing {
// Convert TaskInfo to VideoTask
var videoTask types.VideoTask
if err := json.Unmarshal([]byte(job.TaskInfo), &videoTask); err != nil {
logger.Errorf("failed to unmarshal task info to VideoTask: %v", err)
continue
}
// Type assert task.Params to KeLingVideoParams
paramsMap, ok := videoTask.Params.(map[string]interface{})
if !ok {
continue
}
// Convert map to KeLingVideoParams
paramsBytes, err := json.Marshal(paramsMap)
if err != nil {
continue
}
var params types.KeLingVideoParams
if err := json.Unmarshal(paramsBytes, &params); err != nil {
continue
}
task, err := s.QueryKeLingTask(job.TaskId, job.Channel, params.TaskType)
if err != nil {
logger.Errorf("query task with error: %v", err)
// 更新任务信息
s.db.Model(&model.VideoJob{Id: job.Id}).UpdateColumns(map[string]interface{}{
"progress": service.FailTaskProgress, // 102 表示资源未下载完成,
"err_msg": err.Error(),
"cover_url": "/images/failed.jpg",
})
continue
}
logger.Debugf("task: %+v", task)
if task.TaskStatus == "succeed" { // 更新任务信息
data := map[string]interface{}{
"progress": 102, // 102 表示资源未下载完成,
"water_url": task.TaskResult.Videos[0].URL,
"raw_data": utils.JsonEncode(task),
"prompt_ext": job.Prompt,
"cover_url": "",
}
if len(task.TaskResult.Videos) > 0 {
data["video_url"] = task.TaskResult.Videos[0].URL
}
err = s.db.Model(&model.VideoJob{Id: job.Id}).UpdateColumns(data).Error
if err != nil {
logger.Errorf("更新数据库失败:%v", err)
continue
}
} else if task.TaskStatus == "failed" {
// 更新任务信息
s.db.Model(&model.VideoJob{Id: job.Id}).UpdateColumns(map[string]interface{}{
"progress": service.FailTaskProgress,
"err_msg": task.TaskStatusMsg,
"cover_url": "/images/failed.jpg",
})
}
}
}
// 找出失败的任务,并恢复其扣减算力
s.db.Where("progress", service.FailTaskProgress).Where("power > ?", 0).Find(&jobs)
for _, job := range jobs {
err := s.userService.IncreasePower(job.UserId, job.Power, model.PowerLog{
Type: types.PowerRefund,
Model: job.Type,
Remark: fmt.Sprintf("%s 任务失败,退回算力。任务ID%sErr:%s", job.Type, job.TaskId, job.ErrMsg),
})
if err != nil {
continue
}
// 更新任务状态
s.db.Model(&job).UpdateColumn("power", 0)
}
time.Sleep(time.Second * 10)
}
}()
}
type LumaTaskVo struct {
Id string `json:"id"`
Liked interface{} `json:"liked"`
State string `json:"state"`
Video struct {
Url string `json:"url"`
Width int `json:"width"`
Height int `json:"height"`
Thumbnail string `json:"thumbnail"`
DownloadUrl string `json:"download_url"`
} `json:"video"`
Prompt string `json:"prompt"`
UserId string `json:"user_id"`
BatchId string `json:"batch_id"`
Thumbnail struct {
Url string `json:"url"`
Width int `json:"width"`
Height int `json:"height"`
} `json:"thumbnail"`
VideoRaw struct {
Url string `json:"url"`
Width int `json:"width"`
Height int `json:"height"`
} `json:"video_raw"`
CreatedAt string `json:"created_at"`
LastFrame struct {
Url string `json:"url"`
Width int `json:"width"`
Height int `json:"height"`
} `json:"last_frame"`
}
type LumaRespVo struct {
Id string `json:"id"`
Prompt string `json:"prompt"`
State string `json:"state"`
QueueState interface{} `json:"queue_state"`
CreatedAt string `json:"created_at"`
Video interface{} `json:"video"`
VideoRaw interface{} `json:"video_raw"`
Liked interface{} `json:"liked"`
EstimateWaitSeconds interface{} `json:"estimate_wait_seconds"`
Thumbnail interface{} `json:"thumbnail"`
Channel string `json:"channel,omitempty"`
}
func (s *Service) LumaCreate(task types.VideoTask) (LumaRespVo, error) {
// 读取 API KEY
var apiKey model.ApiKey
session := s.db.Session(&gorm.Session{}).Where("type", "luma").Where("enabled", true)
if task.Channel != "" {
session = session.Where("api_url", task.Channel)
}
tx := session.Order("last_used_at DESC").First(&apiKey)
if tx.Error != nil {
return LumaRespVo{}, errors.New("no available API KEY for Luma")
}
// Type assert task.Params to LumaVideoParams
paramsMap, ok := task.Params.(map[string]interface{})
if !ok {
return LumaRespVo{}, errors.New("invalid params type for Luma video task")
}
// Convert map to LumaVideoParams
paramsBytes, err := json.Marshal(paramsMap)
if err != nil {
return LumaRespVo{}, fmt.Errorf("failed to marshal params: %v", err)
}
var params types.LumaVideoParams
if err := json.Unmarshal(paramsBytes, &params); err != nil {
return LumaRespVo{}, fmt.Errorf("failed to unmarshal params: %v", err)
}
reqBody := map[string]interface{}{
"user_prompt": task.Prompt,
"expand_prompt": params.PromptOptimize,
"loop": params.Loop,
"image_url": params.StartImgURL, // 图生视频
"image_end_url": params.EndImgURL, // 图生视频
}
var res LumaRespVo
apiURL := fmt.Sprintf("%s/luma/generations", apiKey.ApiURL)
logger.Debugf("API URL: %s, request body: %+v", apiURL, reqBody)
r, err := req.C().R().
SetHeader("Authorization", "Bearer "+apiKey.Value).
SetBody(reqBody).
Post(apiURL)
if err != nil {
return LumaRespVo{}, fmt.Errorf("请求 API 出错:%v", err)
}
if r.StatusCode != 200 && r.StatusCode != 201 {
return LumaRespVo{}, fmt.Errorf("请求 API 出错:%d, %s", r.StatusCode, r.String())
}
body, _ := io.ReadAll(r.Body)
err = json.Unmarshal(body, &res)
if err != nil {
return LumaRespVo{}, fmt.Errorf("解析API数据失败:%v, %s", err, string(body))
}
// update the last_use_at for api key
apiKey.LastUsedAt = time.Now().Unix()
session.Updates(&apiKey)
res.Channel = apiKey.ApiURL
return res, nil
}
func (s *Service) QueryLumaTask(taskId string, channel string) (LumaTaskVo, error) {
// 读取 API KEY
var apiKey model.ApiKey
err := s.db.Session(&gorm.Session{}).Where("type", "luma").
Where("api_url", channel).
Where("enabled", true).
Order("last_used_at DESC").First(&apiKey).Error
if err != nil {
return LumaTaskVo{}, errors.New("no available API KEY for Luma")
}
apiURL := fmt.Sprintf("%s/luma/generations/%s", apiKey.ApiURL, taskId)
var res LumaTaskVo
r, err := req.C().R().SetHeader("Authorization", "Bearer "+apiKey.Value).Get(apiURL)
if err != nil {
return LumaTaskVo{}, fmt.Errorf("请求 API 失败:%v", err)
}
defer r.Body.Close()
if r.StatusCode != 200 {
return LumaTaskVo{}, fmt.Errorf("API 返回失败:%v", r.String())
}
body, _ := io.ReadAll(r.Body)
err = json.Unmarshal(body, &res)
if err != nil {
return LumaTaskVo{}, fmt.Errorf("解析API数据失败:%v, %s", err, string(body))
}
return res, nil
}
type KeLingRespVo struct {
Code int `json:"code"`
Message string `json:"message"`
RequestID string `json:"request_id"`
Data struct {
TaskID string `json:"task_id"`
TaskStatus string `json:"task_status"`
CreatedAt int64 `json:"created_at"`
UpdatedAt int64 `json:"updated_at"`
} `json:"data"`
Channel string `json:"channel,omitempty"`
}
func (s *Service) KeLingCreate(task types.VideoTask) (KeLingRespVo, error) {
var apiKey model.ApiKey
session := s.db.Session(&gorm.Session{}).Where("type", "keling").Where("enabled", true)
if task.Channel != "" {
session = session.Where("api_url", task.Channel)
}
tx := session.Order("last_used_at DESC").First(&apiKey)
if tx.Error != nil {
return KeLingRespVo{}, errors.New("no available API KEY for keling")
}
// Type assert task.Params to KeLingVideoParams
paramsMap, ok := task.Params.(map[string]interface{})
if !ok {
return KeLingRespVo{}, errors.New("invalid params type for KeLing video task")
}
// Convert map to KeLingVideoParams
paramsBytes, err := json.Marshal(paramsMap)
if err != nil {
return KeLingRespVo{}, fmt.Errorf("failed to marshal params: %v", err)
}
var params types.KeLingVideoParams
if err := json.Unmarshal(paramsBytes, &params); err != nil {
return KeLingRespVo{}, fmt.Errorf("failed to unmarshal params: %v", err)
}
// 2. 构建API请求参数
payload := map[string]interface{}{
"model_name": params.Model,
"prompt": task.Prompt,
"negative_prompt": params.NegPrompt,
"cfg_scale": params.CfgScale,
"mode": params.Mode,
"aspect_ratio": params.AspectRatio,
"duration": params.Duration,
}
// 只有当 CameraControl 的类型不为空时,才处理摄像机控制参数
if params.CameraControl.Type != "" {
cameraControl := map[string]interface{}{
"type": params.CameraControl.Type,
}
// 只有在 simple 类型时才添加 config 参数
if params.CameraControl.Type == "simple" {
cameraControl["config"] = params.CameraControl.Config
}
payload["camera_control"] = cameraControl
}
// 处理图生视频
if params.TaskType == "image2video" {
payload["image"] = params.Image
payload["image_tail"] = params.ImageTail
}
jsonPayload, err := json.Marshal(payload)
if err != nil {
return KeLingRespVo{}, fmt.Errorf("failed to marshal payload: %v", err)
}
// 3. 准备HTTP请求
url := fmt.Sprintf("%s/kling/v1/videos/%s", apiKey.ApiURL, params.TaskType)
req, err := http.NewRequest("POST", url, bytes.NewReader(jsonPayload))
if err != nil {
return KeLingRespVo{}, fmt.Errorf("failed to create request: %v", err)
}
req.Header.Set("Authorization", "Bearer "+apiKey.Value)
req.Header.Set("Content-Type", "application/json")
// 4. 发送请求
client := &http.Client{Timeout: time.Duration(30) * time.Second}
resp, err := client.Do(req)
if err != nil {
return KeLingRespVo{}, fmt.Errorf("failed to send request: %v", err)
}
defer resp.Body.Close()
// 5. 处理响应
body, err := io.ReadAll(resp.Body)
if err != nil {
return KeLingRespVo{}, fmt.Errorf("failed to read response: %v", err)
}
if resp.StatusCode != http.StatusOK {
return KeLingRespVo{}, fmt.Errorf("API error (status %d): %s", resp.StatusCode, string(body))
}
var apiResponse = KeLingRespVo{}
if err := json.Unmarshal(body, &apiResponse); err != nil {
return KeLingRespVo{}, fmt.Errorf("failed to parse response: %v", err)
}
// 设置 API 通道
apiResponse.Channel = apiKey.ApiURL
return apiResponse, nil
}
// VideoCallbackData 表示视频生成任务的回调数据
type VideoCallbackData struct {
TaskID string `json:"task_id"`
TaskStatus string `json:"task_status"`
TaskStatusMsg string `json:"task_status_msg"`
CreatedAt int64 `json:"created_at"`
UpdatedAt int64 `json:"updated_at"`
TaskResult TaskResult `json:"task_result"`
}
type TaskResult struct {
Images []CallBackImageResult `json:"images,omitempty"`
Videos []CallBackVideoResult `json:"videos,omitempty"`
}
type CallBackImageResult struct {
Index int `json:"index"`
URL string `json:"url"`
}
type CallBackVideoResult struct {
ID string `json:"id"`
URL string `json:"url"`
Duration string `json:"duration"`
}
func (s *Service) QueryKeLingTask(taskId string, channel string, action string) (VideoCallbackData, error) {
var apiKey model.ApiKey
err := s.db.Session(&gorm.Session{}).Where("type", "keling").
//Where("api_url", channel).
Where("enabled", true).
Order("last_used_at DESC").First(&apiKey).Error
if err != nil {
return VideoCallbackData{}, errors.New("no available API KEY for keling")
}
url := fmt.Sprintf("%s/kling/v1/videos/%s/%s", apiKey.ApiURL, action, taskId)
req, err := http.NewRequest("GET", url, nil)
if err != nil {
return VideoCallbackData{}, fmt.Errorf("failed to create request: %w", err)
}
req.Header.Set("Authorization", "Bearer "+apiKey.Value)
req.Header.Set("Content-Type", "application/json")
client := &http.Client{}
res, err := client.Do(req)
if err != nil {
return VideoCallbackData{}, fmt.Errorf("failed to execute request: %w", err)
}
defer res.Body.Close()
if res.StatusCode != http.StatusOK {
return VideoCallbackData{}, fmt.Errorf("unexpected status code: %d", res.StatusCode)
}
body, err := io.ReadAll(res.Body)
if err != nil {
return VideoCallbackData{}, fmt.Errorf("failed to read response body: %w", err)
}
var response struct {
Code int `json:"code"`
Message string `json:"message"`
Data VideoCallbackData `json:"data"`
}
if err := json.Unmarshal(body, &response); err != nil {
return VideoCallbackData{}, fmt.Errorf("failed to unmarshal response: %w", err)
}
if response.Code != 0 {
return VideoCallbackData{}, fmt.Errorf("API error: %s", response.Message)
}
return response.Data, nil
}
+333
View File
@@ -0,0 +1,333 @@
package video
// * +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
// * 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 (
"encoding/json"
"fmt"
"geekai/core/types"
"geekai/log"
"geekai/service"
"geekai/service/oss"
"geekai/service/video/adapters"
"geekai/store"
"geekai/store/model"
"geekai/utils"
"time"
"github.com/go-redis/redis/v8"
"gorm.io/gorm"
)
var logger = log.GetLogger()
type Service struct {
db *gorm.DB
uploadManager *oss.UploaderManager
taskQueue *store.RedisQueue
userService *service.UserService
adapters map[string]adapters.VideoAdapter // provider -> adapter
}
func NewService(db *gorm.DB, manager *oss.UploaderManager, redisCli *redis.Client, userService *service.UserService) *Service {
service := &Service{
db: db,
taskQueue: store.NewRedisQueue("Video_Task_Queue", redisCli),
uploadManager: manager,
userService: userService,
adapters: make(map[string]VideoAdapter),
}
// 注册所有适配器
service.registerAdapters()
return service
}
// VideoAdapter 类型别名,指向 adapters.VideoAdapter
type VideoAdapter = adapters.VideoAdapter
// registerAdapters 注册所有视频生成适配器
func (s *Service) registerAdapters() {
// 注册 Veo 适配器
veoAdapter := adapters.NewVeoAdapter(s.db)
s.adapters[veoAdapter.GetProvider()] = veoAdapter
// 注册 Sora 适配器
soraAdapter := adapters.NewSoraAdapter(s.db)
s.adapters[soraAdapter.GetProvider()] = soraAdapter
// 注册 Luma 适配器
lumaAdapter := adapters.NewLumaAdapter(s.db)
s.adapters[lumaAdapter.GetProvider()] = lumaAdapter
// 注册可灵适配器
kelingAdapter := adapters.NewKelingAdapter(s.db)
s.adapters[kelingAdapter.GetProvider()] = kelingAdapter
// 注册 MiniMax 适配器
minimaxAdapter := adapters.NewMiniMaxAdapter(s.db)
s.adapters[minimaxAdapter.GetProvider()] = minimaxAdapter
// 注册 Wan 适配器
wanAdapter := adapters.NewWanAdapter(s.db)
s.adapters[wanAdapter.GetProvider()] = wanAdapter
// 注册 Doubao 适配器
doubaoAdapter := adapters.NewDoubaoAdapter(s.db)
s.adapters[doubaoAdapter.GetProvider()] = doubaoAdapter
}
// getAdapter 获取指定 provider 的适配器
func (s *Service) getAdapter(provider string) (adapters.VideoAdapter, error) {
adapter, ok := s.adapters[provider]
if !ok {
return nil, fmt.Errorf("不支持的视频生成服务提供商: %s", provider)
}
return adapter, nil
}
// getVideoConfig 获取视频配置
func (s *Service) getVideoConfig() (*types.VideoConfig, error) {
return GetVideoConfig(s.db)
}
// CreateTask 统一的创建任务方法
func (s *Service) CreateTask(task types.VideoTask) (adapters.CreateTaskResponse, error) {
// 获取适配器
adapter, err := s.getAdapter(task.Type)
if err != nil {
return adapters.CreateTaskResponse{}, err
}
// 获取视频配置
videoConfig, err := s.getVideoConfig()
if err != nil {
return adapters.CreateTaskResponse{}, err
}
// 调用适配器创建任务
return adapter.CreateTask(task, videoConfig)
}
// QueryTask 统一的查询任务方法
func (s *Service) QueryTask(provider string, taskId string, channel string, modelKey string) (adapters.QueryTaskResponse, error) {
// 获取适配器
adapter, err := s.getAdapter(provider)
if err != nil {
return adapters.QueryTaskResponse{}, err
}
// 获取视频配置
videoConfig, err := s.getVideoConfig()
if err != nil {
return adapters.QueryTaskResponse{}, err
}
// 调用适配器查询任务
return adapter.QueryTask(taskId, channel, videoConfig)
}
func (s *Service) PushTask(task types.VideoTask) {
logger.Infof("[video] push task to queue jobId=%d type=%s", task.Id, task.Type)
if err := s.taskQueue.RPush(task); err != nil {
logger.Errorf("[video] push task to queue failed jobId=%d: %v", task.Id, err)
}
}
func (s *Service) Run() {
// 将数据库中未提交的任务加载到队列
var jobs []model.VideoJob
s.db.Where("task_id", "").Where("progress", 0).Find(&jobs)
for _, v := range jobs {
var task types.VideoTask
err := utils.JsonDecode(v.Params, &task)
if err != nil {
logger.Errorf("decode task info with error: %v", err)
continue
}
task.Id = v.Id
s.PushTask(task)
}
logger.Infof("[video] job consumer started, loaded %d pending jobs from DB", len(jobs))
go func() {
for {
var task types.VideoTask
err := s.taskQueue.LPop(&task)
if err != nil {
logger.Errorf("taking task with error: %v", err)
continue
}
logger.Debugf("[video] submitting task jobId=%d type=%s prompt=%q", task.Id, task.Type, task.Prompt)
r, err := s.CreateTask(task)
if err != nil {
logger.Errorf("[video] submit failed jobId=%d type=%s: %v", task.Id, task.Type, err)
err = s.db.Model(&model.VideoJob{Id: task.Id}).UpdateColumns(map[string]interface{}{
"err_msg": err.Error(),
"status": types.VideoStatusFailed,
}).Error
if err != nil {
logger.Errorf("update task with error: %v", err)
}
continue
}
logger.Infof("[video] submit success jobId=%d type=%s taskId=%s channel=%s", task.Id, task.Type, r.TaskId, r.Channel)
err = s.db.Model(&model.VideoJob{Id: task.Id}).UpdateColumns(map[string]interface{}{
"task_id": r.TaskId,
"channel": r.Channel,
"status": types.VideoStatusPending,
}).Error
if err != nil {
logger.Errorf("update task with error: %v", err)
s.PushTask(task)
}
}
}()
}
func (s *Service) DownloadFiles() {
go func() {
var items []model.VideoJob
logger.Info("[video] download files started")
for {
err := s.db.Where("status", types.VideoStatusDownloading).Find(&items).Error
if err != nil {
logger.Errorf("get downloading tasks with error: %v", err)
continue
}
for _, v := range items {
if v.VideoURL == "" {
continue
}
logger.Infof("try download video: %s", v.VideoURL)
videoURL, err := s.uploadManager.GetUploadHandler().PutUrlFile(v.VideoURL, ".mp4", true)
if err != nil {
logger.Errorf("download video with error: %v", err)
continue
}
logger.Infof("download video success: %s", videoURL)
s.db.Model(&model.VideoJob{Id: v.Id}).UpdateColumns(map[string]any{
"video_url": videoURL,
"status": types.VideoStatusSuccess,
"progress": 100,
})
}
time.Sleep(time.Second * 10)
}
}()
}
// SyncTaskProgress 异步拉取任务
func (s *Service) SyncTaskProgress() {
go func() {
logger.Info("[video] task status poller started")
var jobs []model.VideoJob
for {
res := s.db.Where("status IN ?", []string{types.VideoStatusInProgress, types.VideoStatusPending}).Where("task_id <> ?", "").Find(&jobs)
if res.Error != nil {
continue
}
if len(jobs) > 0 {
logger.Infof("[video] polling task status, in_progress count=%d", len(jobs))
}
for _, job := range jobs {
// 检查任务是否超时(超过 2 小时)
if time.Since(job.CreatedAt) > 2*time.Hour {
logger.Warnf("[video] task timeout jobId=%d taskId=%s created_at=%s", job.Id, job.TaskId, job.CreatedAt.Format(time.RFC3339))
err := s.db.Model(&model.VideoJob{Id: job.Id}).UpdateColumns(map[string]any{
"status": types.VideoStatusFailed,
"err_msg": "任务超时",
}).Error
if err != nil {
logger.Errorf("[video] update timeout task failed jobId=%d: %v", job.Id, err)
}
continue
}
modelKey := ""
var videoTask types.VideoTask
if err := json.Unmarshal([]byte(job.Params), &videoTask); err == nil {
if paramsMap, ok := videoTask.Params.(map[string]any); ok {
if model, ok := paramsMap["model"].(string); ok {
modelKey = model
}
}
}
logger.Debugf("[video] querying task jobId=%d taskId=%s provider=%s", job.Id, job.TaskId, job.Type)
task, err := s.QueryTask(job.Type, job.TaskId, job.Channel, modelKey)
if err != nil {
logger.Errorf("[video] query failed jobId=%d taskId=%s: %v", job.Id, job.TaskId, err)
// 更新任务信息
s.db.Model(&model.VideoJob{Id: job.Id}).UpdateColumns(map[string]any{
"status": types.VideoStatusFailed,
"err_msg": err.Error(),
})
continue
}
logger.Debugf("[video] task status jobId=%d taskId=%s status=%s", job.Id, job.TaskId, task.Status)
logger.Debugf("[video] output=%s", task.Output)
if task.Status == types.VideoStatusSuccess {
data := map[string]any{
"status": types.VideoStatusDownloading,
"progress": 100,
"output": task.Output,
}
if task.VideoURL != "" {
data["video_url"] = task.VideoURL
}
err = s.db.Model(&model.VideoJob{Id: job.Id}).UpdateColumns(data).Error
if err != nil {
logger.Errorf("更新数据库失败:%v", err)
continue
}
logger.Infof("[video] task completed jobId=%d taskId=%s", job.Id, job.TaskId)
} else if task.Status == "failed" {
logger.Warnf("[video] task failed jobId=%d taskId=%s err=%s", job.Id, job.TaskId, task.ErrMsg)
s.db.Model(&model.VideoJob{Id: job.Id}).UpdateColumns(map[string]any{
"status": types.VideoStatusFailed,
"err_msg": task.ErrMsg,
})
} else {
s.db.Model(&model.VideoJob{Id: job.Id}).UpdateColumns(map[string]any{
"status": task.Status,
"progress": task.Progress,
})
continue
}
}
// 找出失败的任务,并恢复其扣减算力
s.db.Select("id", "user_id", "power", "task_id", "err_msg", "type").
Where("status", types.VideoStatusFailed).Where("power > ?", 0).Find(&jobs)
for _, job := range jobs {
err := s.userService.IncreasePower(job.UserId, job.Power, model.PowerLog{
Type: types.PowerRefund,
Model: job.Type,
Remark: fmt.Sprintf("%s 任务失败,退回算力。任务ID%sErr:%s", job.Type, job.TaskId, job.ErrMsg),
})
if err != nil {
continue
}
// 更新任务状态
s.db.Model(&job).UpdateColumn("power", 0)
}
time.Sleep(time.Second * 10)
}
}()
}
+94
View File
@@ -0,0 +1,94 @@
package service
import (
"context"
"encoding/json"
"fmt"
"geekai/core/types"
"geekai/store/model"
"geekai/utils"
"time"
"gorm.io/gorm"
)
// WxGzhService 微信公众号服务
type WxGzhService struct {
config types.WxGzhConfig
DB *gorm.DB
}
func (s *WxGzhService) UpdateConfig(config types.WxGzhConfig) {
s.config = config
}
func (s *WxGzhService) GetConfig() types.WxGzhConfig {
return s.config
}
func (s *WxGzhService) SetConfig(config types.WxGzhConfig) {
s.config = config
}
func NewWxGzhService(config *types.SystemConfig, db *gorm.DB) *WxGzhService {
return &WxGzhService{config: config.WxGzh, DB: db}
}
// GetOpenIDByCode 根据 code 获取 openid 和 access_token
func (s *WxGzhService) GetOpenIDByCode(code string) (string, string, error) {
var config model.Config
s.DB.Where("name", types.ConfigKeyWxGzh).First(&config)
var value map[string]any
err := utils.JsonDecode(config.Value, &value)
url := fmt.Sprintf("https://api.weixin.qq.com/sns/oauth2/access_token?appid=%s&secret=%s&code=%s&grant_type=authorization_code",
value["app_id"], value["secret"], code)
body, status, err := utils.FetchURLBytes(context.Background(), url, "", 30*time.Second, 2, 2<<20)
if err != nil {
return "", "", fmt.Errorf("wx get openid failed: status=%d: %w", status, err)
}
var result map[string]interface{}
err = json.Unmarshal(body, &result)
if err != nil {
return "", "", err
}
if openID, ok := result["openid"].(string); ok {
if accessToken, ok := result["access_token"].(string); ok {
return openID, accessToken, nil
}
}
if errMsg, ok := result["errmsg"].(string); ok {
return "", "", fmt.Errorf("微信 API 错误: %s", errMsg)
}
return "", "", fmt.Errorf("获取 openid 和 access_token 失败: %s", string(body))
}
// GetUserInfo 获取微信用户昵称和头像
func (s *WxGzhService) GetUserInfo(accessToken string, openID string) (map[string]any, error) {
url := fmt.Sprintf("https://api.weixin.qq.com/sns/userinfo?access_token=%s&openid=%s&lang=zh_CN",
accessToken, openID)
body, status, err := utils.FetchURLBytes(context.Background(), url, "", 30*time.Second, 2, 2<<20)
if err != nil {
return nil, fmt.Errorf("wx get userinfo failed: status=%d: %w", status, err)
}
var result map[string]any
err = json.Unmarshal(body, &result)
if err != nil {
return nil, err
}
if errMsg, ok := result["errmsg"].(string); ok {
return nil, fmt.Errorf("微信 API 错误: %s", errMsg)
}
return result, nil
}