初始化
This commit is contained in:
@@ -0,0 +1,575 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"fmt"
|
||||
"math"
|
||||
"math/big"
|
||||
"net/http"
|
||||
"regexp"
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"juhe-factory/api/internal/security"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type AdminData struct {
|
||||
DB *gorm.DB
|
||||
Passwords security.PasswordHasher
|
||||
Encryptor *security.Encryptor
|
||||
HTTPClient *http.Client
|
||||
EncryptionKeyVersion string
|
||||
balanceSyncing atomic.Bool
|
||||
}
|
||||
|
||||
type Page struct {
|
||||
Items []map[string]any `json:"items"`
|
||||
Total int64 `json:"total"`
|
||||
Page int `json:"page"`
|
||||
PageSize int `json:"page_size"`
|
||||
}
|
||||
|
||||
type ResourceSpec struct {
|
||||
Table string
|
||||
Select string
|
||||
SearchColumns []string
|
||||
Fields map[string]string
|
||||
Filters map[string]string
|
||||
Order string
|
||||
SoftDelete bool
|
||||
}
|
||||
|
||||
var resourceSpecs = map[string]ResourceSpec{
|
||||
"styles": {Table: "project_styles", Select: "id,name,image_key,image_url,image_mime,image_size,image_width,image_height,sort_order", Fields: map[string]string{"name": "name", "image_key": "image_key", "image_url": "image_url", "image_mime": "image_mime", "image_size": "image_size", "image_width": "image_width", "image_height": "image_height"}, Order: "sort_order,name", SoftDelete: true},
|
||||
"channels": {Table: "channels", Select: "c.id,c.name,c.channel_type,c.base_url,c.api_key_ciphertext,c.api_key_last4,c.warning_threshold,c.channel_points_per_cny,c.max_concurrency,c.max_user_concurrency,c.enabled,c.created_at,c.updated_at,b.balance,b.currency,b.synced_at", SearchColumns: []string{"c.name", "c.base_url"}, Fields: map[string]string{"name": "name", "channel_type": "channel_type", "base_url": "base_url", "warning_threshold": "warning_threshold", "channel_points_per_cny": "channel_points_per_cny", "max_concurrency": "max_concurrency", "max_user_concurrency": "max_user_concurrency", "enabled": "enabled"}, Filters: map[string]string{"channel_type": "channel_type", "enabled": "enabled"}, Order: "c.created_at DESC", SoftDelete: true},
|
||||
"models": {Table: "models", Select: "id", Fields: map[string]string{}, Order: "created_at DESC", SoftDelete: true},
|
||||
"prompts": {Table: "prompts", Select: "id", Fields: map[string]string{}, Order: "updated_at DESC", SoftDelete: true},
|
||||
}
|
||||
|
||||
func pageArgs(pageRaw, sizeRaw string) (int, int) {
|
||||
page, _ := strconv.Atoi(pageRaw)
|
||||
if page < 1 {
|
||||
page = 1
|
||||
}
|
||||
size, _ := strconv.Atoi(sizeRaw)
|
||||
if size < 1 {
|
||||
size = 20
|
||||
}
|
||||
if size > 100 {
|
||||
size = 100
|
||||
}
|
||||
return page, size
|
||||
}
|
||||
|
||||
func (s *AdminData) ListResource(ctx context.Context, resource, keyword string, filters map[string]string, pageRaw, sizeRaw string) (Page, error) {
|
||||
spec, ok := resourceSpecs[resource]
|
||||
if !ok {
|
||||
return Page{}, errors.New("不支持的资源类型")
|
||||
}
|
||||
page, size := pageArgs(pageRaw, sizeRaw)
|
||||
if resource == "channels" {
|
||||
s.triggerChannelBalanceSync()
|
||||
}
|
||||
query := s.DB.Table(spec.Table)
|
||||
if resource == "channels" {
|
||||
query = s.DB.Table("channels c").Joins("LEFT JOIN LATERAL (SELECT balance,currency,synced_at FROM channel_balance_snapshots WHERE channel_id=c.id ORDER BY synced_at DESC LIMIT 1) b ON true")
|
||||
}
|
||||
if spec.SoftDelete {
|
||||
prefix := strings.Split(spec.Table, " ")[0]
|
||||
if resource == "channels" {
|
||||
prefix = "c"
|
||||
}
|
||||
query = query.Where(prefix + ".deleted_at IS NULL")
|
||||
}
|
||||
if keyword = strings.TrimSpace(keyword); keyword != "" && len(spec.SearchColumns) > 0 {
|
||||
parts := make([]string, len(spec.SearchColumns))
|
||||
args := make([]any, len(spec.SearchColumns))
|
||||
for i, column := range spec.SearchColumns {
|
||||
parts[i] = column + "::text ILIKE ?"
|
||||
args[i] = "%" + keyword + "%"
|
||||
}
|
||||
query = query.Where("("+strings.Join(parts, " OR ")+")", args...)
|
||||
}
|
||||
for key, value := range filters {
|
||||
if column, exists := spec.Filters[key]; exists && value != "" {
|
||||
query = query.Where(column+" = ?", value)
|
||||
}
|
||||
}
|
||||
var total int64
|
||||
if err := query.Count(&total).Error; err != nil {
|
||||
return Page{}, err
|
||||
}
|
||||
items := make([]map[string]any, 0)
|
||||
if err := query.Select(spec.Select).Order(spec.Order).Offset((page - 1) * size).Limit(size).Find(&items).Error; err != nil {
|
||||
return Page{}, err
|
||||
}
|
||||
if resource == "channels" {
|
||||
s.maskChannelAPIKeys(items)
|
||||
}
|
||||
return Page{Items: items, Total: total, Page: page, PageSize: size}, nil
|
||||
}
|
||||
|
||||
func (s *AdminData) maskChannelAPIKeys(items []map[string]any) {
|
||||
for _, item := range items {
|
||||
ciphertext := strings.TrimSpace(fmt.Sprint(item["api_key_ciphertext"]))
|
||||
last4 := strings.TrimSpace(fmt.Sprint(item["api_key_last4"]))
|
||||
delete(item, "api_key_ciphertext")
|
||||
delete(item, "api_key_last4")
|
||||
item["api_key_masked"] = ""
|
||||
if ciphertext != "" && ciphertext != "<nil>" && s.Encryptor != nil {
|
||||
if plain, err := s.Encryptor.Decrypt(ciphertext); err == nil {
|
||||
item["api_key_masked"] = maskAPIKey(plain)
|
||||
continue
|
||||
}
|
||||
}
|
||||
if last4 != "" && last4 != "<nil>" {
|
||||
item["api_key_masked"] = "********" + last4
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func maskAPIKey(value string) string {
|
||||
characters := []rune(value)
|
||||
if len(characters) <= 10 {
|
||||
return strings.Repeat("*", max(len(characters), 8))
|
||||
}
|
||||
return string(characters[:5]) + "********" + string(characters[len(characters)-5:])
|
||||
}
|
||||
|
||||
func (s *AdminData) SaveResource(resource, id string, values map[string]any) (string, error) {
|
||||
spec, ok := resourceSpecs[resource]
|
||||
if !ok {
|
||||
return "", errors.New("不支持的资源类型")
|
||||
}
|
||||
data := map[string]any{}
|
||||
for input, column := range spec.Fields {
|
||||
if value, exists := values[input]; exists {
|
||||
data[column] = value
|
||||
}
|
||||
}
|
||||
if resource == "channels" {
|
||||
if err := s.prepareChannel(values, data); err != nil {
|
||||
return "", err
|
||||
}
|
||||
}
|
||||
if resource == "channels" && id == "" {
|
||||
apiKey, _ := values["api_key"].(string)
|
||||
if strings.TrimSpace(apiKey) == "" {
|
||||
return "", errors.New("API Key 不能为空")
|
||||
}
|
||||
data["channel_type"] = "relay"
|
||||
}
|
||||
if id == "" {
|
||||
newID := uuid.NewString()
|
||||
data["id"] = newID
|
||||
if resource == "styles" {
|
||||
var next int
|
||||
if err := s.DB.Table("project_styles").Select("coalesce(max(sort_order),-1)+1").Scan(&next).Error; err != nil {
|
||||
return "", err
|
||||
}
|
||||
data["sort_order"] = next
|
||||
}
|
||||
if err := s.DB.Table(spec.Table).Create(data).Error; err != nil {
|
||||
return "", err
|
||||
}
|
||||
return newID, nil
|
||||
}
|
||||
if len(data) == 0 {
|
||||
return id, nil
|
||||
}
|
||||
result := s.DB.Table(spec.Table).Where("id = ? AND deleted_at IS NULL", id).Updates(data)
|
||||
if result.Error != nil {
|
||||
return "", result.Error
|
||||
}
|
||||
if result.RowsAffected == 0 {
|
||||
return "", gorm.ErrRecordNotFound
|
||||
}
|
||||
return id, nil
|
||||
}
|
||||
|
||||
func (s *AdminData) ReorderStyles(ids []string) error {
|
||||
if len(ids) == 0 {
|
||||
return errors.New("风格排序不能为空")
|
||||
}
|
||||
seen := make(map[string]struct{}, len(ids))
|
||||
for _, id := range ids {
|
||||
if _, err := uuid.Parse(id); err != nil {
|
||||
return errors.New("风格 ID 无效")
|
||||
}
|
||||
if _, exists := seen[id]; exists {
|
||||
return errors.New("风格排序包含重复项")
|
||||
}
|
||||
seen[id] = struct{}{}
|
||||
}
|
||||
return s.DB.Transaction(func(tx *gorm.DB) error {
|
||||
var count int64
|
||||
if err := tx.Table("project_styles").Where("deleted_at IS NULL").Count(&count).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if int64(len(ids)) != count {
|
||||
return errors.New("风格列表已变化,请刷新后重试")
|
||||
}
|
||||
for position, id := range ids {
|
||||
result := tx.Table("project_styles").Where("id=? AND deleted_at IS NULL", id).Update("sort_order", position)
|
||||
if result.Error != nil {
|
||||
return result.Error
|
||||
}
|
||||
if result.RowsAffected == 0 {
|
||||
return errors.New("风格列表已变化,请刷新后重试")
|
||||
}
|
||||
}
|
||||
return nil
|
||||
})
|
||||
}
|
||||
|
||||
func (s *AdminData) prepareChannel(values, data map[string]any) error {
|
||||
rateText := strings.TrimSpace(fmt.Sprint(values["channel_points_per_cny"]))
|
||||
rate, err := strconv.ParseFloat(rateText, 64)
|
||||
if err != nil || math.IsNaN(rate) || math.IsInf(rate, 0) || rate <= 0 || rate > 1000000000000 {
|
||||
return errors.New("渠道汇率必须大于 0")
|
||||
}
|
||||
data["channel_points_per_cny"] = rateText
|
||||
maxConcurrency, err := positiveInt(values["max_concurrency"], 500)
|
||||
if err != nil || maxConcurrency > 5000 {
|
||||
return errors.New("最大并发数必须为 1 至 5000")
|
||||
}
|
||||
maxUserConcurrency, err := positiveInt(values["max_user_concurrency"], 10)
|
||||
if err != nil || maxUserConcurrency > 500 || maxUserConcurrency > maxConcurrency {
|
||||
return errors.New("单用户最高并发数必须为 1 至 500,且不能超过最大并发数")
|
||||
}
|
||||
data["max_concurrency"] = maxConcurrency
|
||||
data["max_user_concurrency"] = maxUserConcurrency
|
||||
plain, _ := values["api_key"].(string)
|
||||
plain = strings.TrimSpace(plain)
|
||||
if plain == "" {
|
||||
return nil
|
||||
}
|
||||
if s.Encryptor == nil {
|
||||
return errors.New("敏感配置加密未配置")
|
||||
}
|
||||
ciphertext, err := s.Encryptor.Encrypt(plain)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
last4 := plain
|
||||
if len(last4) > 4 {
|
||||
last4 = last4[len(last4)-4:]
|
||||
}
|
||||
data["api_key_ciphertext"] = ciphertext
|
||||
data["api_key_last4"] = last4
|
||||
data["encryption_key_version"] = s.EncryptionKeyVersion
|
||||
return nil
|
||||
}
|
||||
|
||||
func positiveInt(value any, fallback int) (int, error) {
|
||||
if value == nil || fmt.Sprint(value) == "" {
|
||||
return fallback, nil
|
||||
}
|
||||
parsed, err := strconv.Atoi(fmt.Sprint(value))
|
||||
if err != nil || parsed < 1 {
|
||||
return 0, errors.New("必须为正整数")
|
||||
}
|
||||
return parsed, nil
|
||||
}
|
||||
|
||||
func (s *AdminData) ToggleResource(resource, id string, enabled bool) error {
|
||||
spec, ok := resourceSpecs[resource]
|
||||
if !ok {
|
||||
return errors.New("不支持的资源类型")
|
||||
}
|
||||
if resource == "models" && enabled {
|
||||
var channelEnabled bool
|
||||
if err := s.DB.Raw(`SELECT c.enabled FROM models m JOIN channels c ON c.id=m.channel_id
|
||||
WHERE m.id=? AND m.deleted_at IS NULL AND c.deleted_at IS NULL`, id).Scan(&channelEnabled).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if !channelEnabled {
|
||||
return errors.New("所属渠道已禁用,无法启用该模型")
|
||||
}
|
||||
}
|
||||
if resource == "channels" {
|
||||
return s.DB.Transaction(func(tx *gorm.DB) error {
|
||||
result := tx.Table("channels").Where("id = ? AND deleted_at IS NULL", id).Update("enabled", enabled)
|
||||
if result.Error != nil {
|
||||
return result.Error
|
||||
}
|
||||
if result.RowsAffected == 0 {
|
||||
return gorm.ErrRecordNotFound
|
||||
}
|
||||
return tx.Table("models").Where("channel_id = ? AND deleted_at IS NULL", id).Update("enabled", enabled).Error
|
||||
})
|
||||
}
|
||||
result := s.DB.Table(spec.Table).Where("id = ? AND deleted_at IS NULL", id).Update("enabled", enabled)
|
||||
if result.Error != nil {
|
||||
return result.Error
|
||||
}
|
||||
if result.RowsAffected == 0 {
|
||||
return gorm.ErrRecordNotFound
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *AdminData) DeleteResource(resource, id string) error {
|
||||
spec, ok := resourceSpecs[resource]
|
||||
if !ok || !spec.SoftDelete {
|
||||
return errors.New("不支持删除该资源")
|
||||
}
|
||||
result := s.DB.Table(spec.Table).Where("id = ? AND deleted_at IS NULL", id).Update("deleted_at", time.Now())
|
||||
if result.Error != nil {
|
||||
return result.Error
|
||||
}
|
||||
if result.RowsAffected == 0 {
|
||||
return gorm.ErrRecordNotFound
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *AdminData) ListUsers(keyword, enabled, pageRaw, sizeRaw string) (Page, error) {
|
||||
page, size := pageArgs(pageRaw, sizeRaw)
|
||||
query := s.DB.Table("web_users u").Where("u.deleted_at IS NULL")
|
||||
if keyword = strings.TrimSpace(keyword); keyword != "" {
|
||||
query = query.Where("u.account::text ILIKE ? OR u.username::text ILIKE ? OR u.uid = ?", "%"+keyword+"%", "%"+keyword+"%", keyword)
|
||||
}
|
||||
if enabled != "" {
|
||||
query = query.Where("u.enabled = ?", enabled)
|
||||
}
|
||||
var total int64
|
||||
if err := query.Count(&total).Error; err != nil {
|
||||
return Page{}, err
|
||||
}
|
||||
items := make([]map[string]any, 0)
|
||||
err := query.Select(`u.id,u.account,u.username,u.uid,u.point_balance,u.daily_limit,u.enabled,u.last_online_at,u.created_at,
|
||||
coalesce((SELECT sum(-l.change_amount) FROM point_ledger l WHERE l.user_id=u.id AND l.change_amount<0 AND l.created_at>=CURRENT_DATE AND l.created_at<CURRENT_DATE+INTERVAL '1 day'),0) AS today_consumption`).
|
||||
Order("u.created_at DESC").Offset((page - 1) * size).Limit(size).Find(&items).Error
|
||||
return Page{Items: items, Total: total, Page: page, PageSize: size}, err
|
||||
}
|
||||
|
||||
func (s *AdminData) CreateUser(account, password string, dailyLimit any) (map[string]any, error) {
|
||||
account = strings.TrimSpace(account)
|
||||
if account == "" {
|
||||
return nil, errors.New("账号不能为空")
|
||||
}
|
||||
items, err := s.createUsers([]string{account}, password, dailyLimit)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return items[0], nil
|
||||
}
|
||||
|
||||
func (s *AdminData) BatchCreateUsers(prefix, startSequence, endSequence, password string, dailyLimit any) ([]map[string]any, error) {
|
||||
accounts, err := buildBatchAccounts(prefix, startSequence, endSequence)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return s.createUsers(accounts, password, dailyLimit)
|
||||
}
|
||||
|
||||
func (s *AdminData) createUsers(accounts []string, password string, dailyLimit any) ([]map[string]any, error) {
|
||||
hash, err := s.Passwords.Hash(password)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
created := make([]map[string]any, 0, len(accounts))
|
||||
err = s.DB.Transaction(func(tx *gorm.DB) error {
|
||||
var existing string
|
||||
if err := tx.Table("web_users").Select("account").Where("account IN ?", accounts).Limit(1).Scan(&existing).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if existing != "" {
|
||||
return fmt.Errorf("账号 %s 已存在", existing)
|
||||
}
|
||||
for _, account := range accounts {
|
||||
user, err := createUserWithAccount(tx, account, hash, dailyLimit)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
created = append(created, user)
|
||||
}
|
||||
return nil
|
||||
})
|
||||
return created, err
|
||||
}
|
||||
|
||||
const usernameAlphabet = "abcdefghijklmnopqrstuvwxyz0123456789"
|
||||
|
||||
var sequencePattern = regexp.MustCompile(`^[0-9]+$`)
|
||||
|
||||
func buildBatchAccounts(prefix, startSequence, endSequence string) ([]string, error) {
|
||||
prefix = strings.TrimSpace(prefix)
|
||||
startSequence = strings.TrimSpace(startSequence)
|
||||
endSequence = strings.TrimSpace(endSequence)
|
||||
if prefix == "" {
|
||||
return nil, errors.New("账号前缀不能为空")
|
||||
}
|
||||
if !sequencePattern.MatchString(startSequence) || !sequencePattern.MatchString(endSequence) {
|
||||
return nil, errors.New("起始序号和终止序号必须为数字")
|
||||
}
|
||||
start, err := strconv.Atoi(startSequence)
|
||||
if err != nil {
|
||||
return nil, errors.New("起始序号无效")
|
||||
}
|
||||
end, err := strconv.Atoi(endSequence)
|
||||
if err != nil {
|
||||
return nil, errors.New("终止序号无效")
|
||||
}
|
||||
if start > end {
|
||||
return nil, errors.New("起始序号不能大于终止序号")
|
||||
}
|
||||
count := end - start + 1
|
||||
if count > 50 {
|
||||
return nil, errors.New("每次最多批量创建 50 个用户")
|
||||
}
|
||||
width := max(len(startSequence), len(endSequence))
|
||||
accounts := make([]string, 0, count)
|
||||
for sequence := start; sequence <= end; sequence++ {
|
||||
accounts = append(accounts, fmt.Sprintf("%s-%0*d", prefix, width, sequence))
|
||||
}
|
||||
return accounts, nil
|
||||
}
|
||||
|
||||
func createUserWithAccount(tx *gorm.DB, account, passwordHash string, dailyLimit any) (map[string]any, error) {
|
||||
for range 20 {
|
||||
username, err := randomUsername()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
uid, err := newUID(tx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
id := uuid.NewString()
|
||||
result := tx.Exec(`INSERT INTO web_users(id,uid,username,account,password_hash,daily_limit,enabled)
|
||||
VALUES(?,?,?,?,?,?,true) ON CONFLICT DO NOTHING`, id, uid, username, account, passwordHash, dailyLimit)
|
||||
if result.Error != nil {
|
||||
return nil, result.Error
|
||||
}
|
||||
if result.RowsAffected == 1 {
|
||||
return map[string]any{"id": id, "uid": uid, "account": account, "username": username}, nil
|
||||
}
|
||||
var accountExists int64
|
||||
if err := tx.Table("web_users").Where("account = ?", account).Count(&accountExists).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if accountExists > 0 {
|
||||
return nil, fmt.Errorf("账号 %s 已存在", account)
|
||||
}
|
||||
}
|
||||
return nil, errors.New("用户名生成失败,请重试")
|
||||
}
|
||||
|
||||
func randomUsername() (string, error) {
|
||||
lengthOffset, err := rand.Int(rand.Reader, big.NewInt(8))
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
username := make([]byte, 5+lengthOffset.Int64())
|
||||
for i := range username {
|
||||
index, err := rand.Int(rand.Reader, big.NewInt(int64(len(usernameAlphabet))))
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
username[i] = usernameAlphabet[index.Int64()]
|
||||
}
|
||||
return string(username), nil
|
||||
}
|
||||
|
||||
func newUID(tx *gorm.DB) (string, error) {
|
||||
for i := 0; i < 20; i++ {
|
||||
n, err := rand.Int(rand.Reader, big.NewInt(90000000))
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
uid := fmt.Sprintf("%08d", n.Int64()+10000000)
|
||||
var count int64
|
||||
if err := tx.Table("web_users").Where("uid=?", uid).Count(&count).Error; err != nil {
|
||||
return "", err
|
||||
}
|
||||
if count == 0 {
|
||||
return uid, nil
|
||||
}
|
||||
}
|
||||
return "", errors.New("UID 生成失败,请重试")
|
||||
}
|
||||
|
||||
func (s *AdminData) UpdateUsers(ids []string, updates map[string]any) error {
|
||||
if len(ids) == 0 {
|
||||
return errors.New("请选择用户")
|
||||
}
|
||||
allowed := map[string]any{}
|
||||
if value, ok := updates["daily_limit"]; ok {
|
||||
allowed["daily_limit"] = value
|
||||
}
|
||||
if value, ok := updates["enabled"]; ok {
|
||||
allowed["enabled"] = value
|
||||
allowed["session_version"] = gorm.Expr("session_version + 1")
|
||||
}
|
||||
if len(allowed) == 0 {
|
||||
return errors.New("没有可更新的字段")
|
||||
}
|
||||
return s.DB.Table("web_users").Where("id IN ? AND deleted_at IS NULL", ids).Updates(allowed).Error
|
||||
}
|
||||
|
||||
func (s *AdminData) DeleteUsers(ids []string) error {
|
||||
if len(ids) == 0 {
|
||||
return errors.New("请选择用户")
|
||||
}
|
||||
return s.DB.Table("web_users").Where("id IN ? AND deleted_at IS NULL", ids).Updates(map[string]any{"deleted_at": time.Now(), "enabled": false, "session_version": gorm.Expr("session_version + 1")}).Error
|
||||
}
|
||||
|
||||
var placeholderPattern = regexp.MustCompile(`\{\{\s*([a-zA-Z_][a-zA-Z0-9_]*)\s*\}\}`)
|
||||
|
||||
func ValidatePromptVariables(content string, variables []string) error {
|
||||
allowed := map[string]bool{}
|
||||
for _, v := range variables {
|
||||
allowed[v] = true
|
||||
}
|
||||
missing := []string{}
|
||||
for _, m := range placeholderPattern.FindAllStringSubmatch(content, -1) {
|
||||
if !allowed[m[1]] {
|
||||
missing = append(missing, m[1])
|
||||
}
|
||||
}
|
||||
sort.Strings(missing)
|
||||
if len(missing) > 0 {
|
||||
return fmt.Errorf("正文使用了未声明变量:%s", strings.Join(missing, "、"))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func GenerateRedemptionCode() (plain, hash, mask string, err error) {
|
||||
const alphabet = "0123456789ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz"
|
||||
const firstAlphabet = "123456789ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz"
|
||||
raw := make([]byte, 30)
|
||||
for i := range raw {
|
||||
chars := alphabet
|
||||
if i == 0 {
|
||||
chars = firstAlphabet
|
||||
}
|
||||
index, randomErr := rand.Int(rand.Reader, big.NewInt(int64(len(chars))))
|
||||
if randomErr != nil {
|
||||
err = randomErr
|
||||
return
|
||||
}
|
||||
raw[i] = chars[index.Int64()]
|
||||
}
|
||||
normalized := string(raw)
|
||||
groups := make([]string, 0, 5)
|
||||
for start := 0; start < len(raw); start += 6 {
|
||||
groups = append(groups, string(raw[start:start+6]))
|
||||
}
|
||||
plain = strings.Join(groups, "-")
|
||||
sum := sha256.Sum256([]byte(normalized))
|
||||
hash = hex.EncodeToString(sum[:])
|
||||
mask = normalized[:3] + "***-******-******-******-" + normalized[len(normalized)-3:]
|
||||
return
|
||||
}
|
||||
@@ -0,0 +1,77 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"reflect"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestRandomUsernameLengthAndCharacters(t *testing.T) {
|
||||
for range 200 {
|
||||
username, err := randomUsername()
|
||||
if err != nil {
|
||||
t.Fatalf("randomUsername returned error: %v", err)
|
||||
}
|
||||
if len(username) < 5 || len(username) > 12 {
|
||||
t.Fatalf("username length must be 5 to 12, got %q", username)
|
||||
}
|
||||
for _, character := range username {
|
||||
if !strings.ContainsRune(usernameAlphabet, character) {
|
||||
t.Fatalf("username contains unsupported character: %q", username)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildBatchAccounts(t *testing.T) {
|
||||
accounts, err := buildBatchAccounts("RR", "001", "003")
|
||||
if err != nil {
|
||||
t.Fatalf("buildBatchAccounts returned error: %v", err)
|
||||
}
|
||||
want := []string{"RR-001", "RR-002", "RR-003"}
|
||||
if !reflect.DeepEqual(accounts, want) {
|
||||
t.Fatalf("accounts = %#v, want %#v", accounts, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildBatchAccountsUsesWidestSequence(t *testing.T) {
|
||||
accounts, err := buildBatchAccounts("A", "1", "003")
|
||||
if err != nil {
|
||||
t.Fatalf("buildBatchAccounts returned error: %v", err)
|
||||
}
|
||||
want := []string{"A-001", "A-002", "A-003"}
|
||||
if !reflect.DeepEqual(accounts, want) {
|
||||
t.Fatalf("accounts = %#v, want %#v", accounts, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildBatchAccountsAllowsFiftyUsers(t *testing.T) {
|
||||
accounts, err := buildBatchAccounts("RR", "001", "050")
|
||||
if err != nil {
|
||||
t.Fatalf("buildBatchAccounts returned error: %v", err)
|
||||
}
|
||||
if len(accounts) != 50 || accounts[0] != "RR-001" || accounts[49] != "RR-050" {
|
||||
t.Fatalf("unexpected account range: first=%q last=%q count=%d", accounts[0], accounts[len(accounts)-1], len(accounts))
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildBatchAccountsValidation(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
prefix string
|
||||
start string
|
||||
end string
|
||||
}{
|
||||
{name: "empty prefix", start: "001", end: "002"},
|
||||
{name: "non numeric start", prefix: "RR", start: "A01", end: "002"},
|
||||
{name: "start after end", prefix: "RR", start: "003", end: "002"},
|
||||
{name: "more than fifty", prefix: "RR", start: "001", end: "051"},
|
||||
}
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
if _, err := buildBatchAccounts(test.prefix, test.start, test.end); err == nil {
|
||||
t.Fatal("expected validation error")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,179 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"juhe-factory/api/internal/model"
|
||||
"juhe-factory/api/internal/security"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
var ErrUnauthorized = errors.New("账号或密码错误")
|
||||
var ErrReauthFailed = errors.New("当前管理员密码错误")
|
||||
var ErrAdminPasswordConfirmation = errors.New("两次输入的新密码不一致")
|
||||
var ErrAdminPasswordLength = errors.New("新密码长度必须为 8~20 位")
|
||||
|
||||
type Auth struct {
|
||||
db *gorm.DB
|
||||
passwords security.PasswordHasher
|
||||
tokens security.TokenService
|
||||
}
|
||||
|
||||
type TokenPair struct {
|
||||
AccessToken string `json:"access_token"`
|
||||
RefreshToken string `json:"refresh_token"`
|
||||
ExpiresAt time.Time `json:"expires_at"`
|
||||
Admin map[string]any `json:"admin"`
|
||||
}
|
||||
|
||||
func NewAuth(db *gorm.DB, passwords security.PasswordHasher, tokens security.TokenService) *Auth {
|
||||
return &Auth{db: db, passwords: passwords, tokens: tokens}
|
||||
}
|
||||
|
||||
func (s *Auth) Bootstrap(username, password string) error {
|
||||
if strings.TrimSpace(password) == "" {
|
||||
return nil
|
||||
}
|
||||
username = strings.TrimSpace(username)
|
||||
var admin model.AdminUser
|
||||
err := s.db.Where("username = ? AND deleted_at IS NULL", username).First(&admin).Error
|
||||
if err == nil {
|
||||
if s.passwords.Verify(admin.PasswordHash, password) {
|
||||
return nil
|
||||
}
|
||||
hash, hashErr := s.passwords.Hash(password)
|
||||
if hashErr != nil {
|
||||
return hashErr
|
||||
}
|
||||
return s.db.Model(&admin).Update("password_hash", hash).Error
|
||||
}
|
||||
if !errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return err
|
||||
}
|
||||
var count int64
|
||||
if err := s.db.Model(&model.AdminUser{}).Count(&count).Error; err != nil || count > 0 {
|
||||
return err
|
||||
}
|
||||
hash, err := s.passwords.Hash(password)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return s.db.Create(&model.AdminUser{Username: username, PasswordHash: hash, Enabled: true}).Error
|
||||
}
|
||||
|
||||
func (s *Auth) Login(username, password string) (*TokenPair, error) {
|
||||
var admin model.AdminUser
|
||||
if err := s.db.Where("username = ? AND deleted_at IS NULL", strings.TrimSpace(username)).First(&admin).Error; err != nil || !admin.Enabled || !s.passwords.Verify(admin.PasswordHash, password) {
|
||||
return nil, ErrUnauthorized
|
||||
}
|
||||
return s.issue(&admin, true)
|
||||
}
|
||||
|
||||
func (s *Auth) Refresh(refreshToken string) (*TokenPair, error) {
|
||||
hash := security.HashToken(refreshToken)
|
||||
var record model.AdminRefreshToken
|
||||
if err := s.db.Where("token_hash = ? AND revoked_at IS NULL AND expires_at > ?", hash, time.Now()).First(&record).Error; err != nil {
|
||||
return nil, ErrUnauthorized
|
||||
}
|
||||
var admin model.AdminUser
|
||||
if err := s.db.First(&admin, "id = ? AND enabled = true", record.AdminID).Error; err != nil {
|
||||
return nil, ErrUnauthorized
|
||||
}
|
||||
now := time.Now()
|
||||
if err := s.db.Model(&record).Update("revoked_at", now).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return s.issue(&admin, false)
|
||||
}
|
||||
|
||||
func (s *Auth) Logout(refreshToken string) error {
|
||||
if refreshToken == "" {
|
||||
return nil
|
||||
}
|
||||
now := time.Now()
|
||||
return s.db.Model(&model.AdminRefreshToken{}).Where("token_hash = ? AND revoked_at IS NULL", security.HashToken(refreshToken)).Update("revoked_at", now).Error
|
||||
}
|
||||
|
||||
func (s *Auth) Authenticate(raw string) (*model.AdminUser, error) {
|
||||
claims, err := s.tokens.ParseAccessToken(raw)
|
||||
if err != nil {
|
||||
return nil, ErrUnauthorized
|
||||
}
|
||||
id, err := uuid.Parse(claims.AdminID)
|
||||
if err != nil {
|
||||
return nil, ErrUnauthorized
|
||||
}
|
||||
var admin model.AdminUser
|
||||
if err := s.db.First(&admin, "id = ? AND enabled = true", id).Error; err != nil {
|
||||
return nil, ErrUnauthorized
|
||||
}
|
||||
return &admin, nil
|
||||
}
|
||||
|
||||
func (s *Auth) ChangePassword(adminID uuid.UUID, currentPassword, newPassword, confirmPassword string) error {
|
||||
if newPassword != confirmPassword {
|
||||
return ErrAdminPasswordConfirmation
|
||||
}
|
||||
if len(newPassword) < 8 || len(newPassword) > 20 {
|
||||
return ErrAdminPasswordLength
|
||||
}
|
||||
var admin model.AdminUser
|
||||
if err := s.db.Where("id = ? AND enabled = true AND deleted_at IS NULL", adminID).First(&admin).Error; err != nil {
|
||||
return ErrUnauthorized
|
||||
}
|
||||
if !s.passwords.Verify(admin.PasswordHash, currentPassword) {
|
||||
return ErrReauthFailed
|
||||
}
|
||||
hash, err := s.passwords.Hash(newPassword)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
now := time.Now()
|
||||
return s.db.Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Model(&admin).Update("password_hash", hash).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.Model(&model.AdminRefreshToken{}).
|
||||
Where("admin_id = ? AND revoked_at IS NULL", adminID).
|
||||
Update("revoked_at", now).Error
|
||||
})
|
||||
}
|
||||
|
||||
func (s *Auth) Audit(admin *model.AdminUser, action, resource, resourceID, reason, ip, traceID string, detail any) error {
|
||||
return WriteAudit(s.db, admin, action, resource, resourceID, reason, ip, traceID, detail)
|
||||
}
|
||||
|
||||
func (s *Auth) issue(admin *model.AdminUser, updateLogin bool) (*TokenPair, error) {
|
||||
access, expiresAt, err := s.tokens.NewAccessToken(admin.ID.String(), admin.Username)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
refresh, hash, refreshExpires, err := s.tokens.NewRefreshToken()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
record := model.AdminRefreshToken{AdminID: admin.ID, TokenHash: hash, ExpiresAt: refreshExpires}
|
||||
if err := s.db.Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Create(&record).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if updateLogin {
|
||||
return tx.Model(admin).Update("last_login_at", time.Now()).Error
|
||||
}
|
||||
return nil
|
||||
}); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &TokenPair{AccessToken: access, RefreshToken: refresh, ExpiresAt: expiresAt, Admin: map[string]any{"id": admin.ID, "username": admin.Username}}, nil
|
||||
}
|
||||
|
||||
func WriteAudit(db *gorm.DB, admin *model.AdminUser, action, resource, resourceID, reason, ip, traceID string, detail any) error {
|
||||
data, _ := json.Marshal(detail)
|
||||
log := model.AdminAuditLog{AdminID: &admin.ID, AdminUsername: admin.Username, Action: action, ResourceType: resource, ResourceID: resourceID, Reason: reason, Detail: string(data), IPAddress: ip, TraceID: traceID}
|
||||
return db.Create(&log).Error
|
||||
}
|
||||
@@ -0,0 +1,78 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"juhe-factory/api/internal/provider/apimart"
|
||||
|
||||
"github.com/google/uuid"
|
||||
)
|
||||
|
||||
type balanceChannel struct {
|
||||
ID uuid.UUID
|
||||
Name string
|
||||
BaseURL string
|
||||
APIKeyCiphertext string
|
||||
}
|
||||
|
||||
func (s *AdminData) triggerChannelBalanceSync() {
|
||||
if !s.balanceSyncing.CompareAndSwap(false, true) {
|
||||
return
|
||||
}
|
||||
go func() {
|
||||
defer s.balanceSyncing.Store(false)
|
||||
s.syncChannelBalances(context.Background())
|
||||
}()
|
||||
}
|
||||
|
||||
func (s *AdminData) syncChannelBalances(ctx context.Context) {
|
||||
if s.Encryptor == nil {
|
||||
return
|
||||
}
|
||||
var channels []balanceChannel
|
||||
if err := s.DB.Table("channels").Select("id,name,base_url,api_key_ciphertext").
|
||||
Where("deleted_at IS NULL AND api_key_ciphertext IS NOT NULL AND api_key_ciphertext<>''").
|
||||
Find(&channels).Error; err != nil {
|
||||
slog.Warn("读取待同步渠道余额失败", "error", err)
|
||||
return
|
||||
}
|
||||
|
||||
syncCtx, cancel := context.WithTimeout(ctx, 10*time.Second)
|
||||
defer cancel()
|
||||
client := s.HTTPClient
|
||||
if client == nil {
|
||||
client = http.DefaultClient
|
||||
}
|
||||
|
||||
var wait sync.WaitGroup
|
||||
for _, channel := range channels {
|
||||
channel := channel
|
||||
wait.Add(1)
|
||||
go func() {
|
||||
defer wait.Done()
|
||||
apiKey, err := s.Encryptor.Decrypt(channel.APIKeyCiphertext)
|
||||
if err != nil {
|
||||
slog.Warn("解密渠道 API Key 失败", "channel", channel.Name, "error", err)
|
||||
return
|
||||
}
|
||||
balance, err := apimart.NewClient(client).Balance(syncCtx, channel.BaseURL, apiKey)
|
||||
if err != nil {
|
||||
slog.Warn("同步渠道余额失败", "channel", channel.Name, "error", err)
|
||||
return
|
||||
}
|
||||
remainBalance := balance.RemainBalance
|
||||
if balance.UnlimitedQuota {
|
||||
remainBalance = -1
|
||||
}
|
||||
if err := s.DB.Exec(`INSERT INTO channel_balance_snapshots(channel_id,balance,currency,synced_at)
|
||||
VALUES(?,?,?,?)`, channel.ID, remainBalance, "USD", time.Now()).Error; err != nil {
|
||||
slog.Warn("保存渠道余额快照失败", "channel", channel.Name, "error", err)
|
||||
}
|
||||
}()
|
||||
}
|
||||
wait.Wait()
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,22 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"sync"
|
||||
"testing"
|
||||
|
||||
"gorm.io/gorm/schema"
|
||||
)
|
||||
|
||||
func TestAnalysisModelConfigIgnoresComputedPricing(t *testing.T) {
|
||||
_, err := schema.Parse(&analysisModelConfig{}, &sync.Map{}, schema.NamingStrategy{})
|
||||
if err != nil {
|
||||
t.Fatalf("parse analysis model config: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestScriptAnalysisModelConfigIgnoresComputedPricing(t *testing.T) {
|
||||
_, err := schema.Parse(&scriptAnalysisModelConfig{}, &sync.Map{}, schema.NamingStrategy{})
|
||||
if err != nil {
|
||||
t.Fatalf("parse script analysis model config: %v", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,113 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type deletionMediaRow struct {
|
||||
ID uuid.UUID
|
||||
ObjectKey string
|
||||
}
|
||||
|
||||
type deletionMediaSet struct {
|
||||
rows map[uuid.UUID]string
|
||||
}
|
||||
|
||||
func newDeletionMediaSet() *deletionMediaSet {
|
||||
return &deletionMediaSet{rows: make(map[uuid.UUID]string)}
|
||||
}
|
||||
|
||||
func (set *deletionMediaSet) addQuery(query *gorm.DB) error {
|
||||
var rows []deletionMediaRow
|
||||
if err := query.Scan(&rows).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
for _, row := range rows {
|
||||
if row.ID != uuid.Nil {
|
||||
set.rows[row.ID] = row.ObjectKey
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (set *deletionMediaSet) objectKeys() []string {
|
||||
keys := make([]string, 0, len(set.rows))
|
||||
seen := make(map[string]struct{}, len(set.rows))
|
||||
for _, key := range set.rows {
|
||||
if key == "" {
|
||||
continue
|
||||
}
|
||||
if _, exists := seen[key]; exists {
|
||||
continue
|
||||
}
|
||||
seen[key] = struct{}{}
|
||||
keys = append(keys, key)
|
||||
}
|
||||
return keys
|
||||
}
|
||||
|
||||
func (set *deletionMediaSet) deleteRows(tx *gorm.DB) error {
|
||||
if len(set.rows) == 0 {
|
||||
return nil
|
||||
}
|
||||
ids := make([]uuid.UUID, 0, len(set.rows))
|
||||
for id := range set.rows {
|
||||
ids = append(ids, id)
|
||||
}
|
||||
if err := tx.Exec("DELETE FROM channel_asset_cache WHERE media_asset_id IN ?", ids).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.Exec("DELETE FROM media_assets WHERE id IN ?", ids).Error
|
||||
}
|
||||
|
||||
func collectMediaByObjectKey(tx *gorm.DB, set *deletionMediaSet, pattern string) error {
|
||||
return set.addQuery(tx.Table("media_assets").Select("id,object_key").Where("object_key LIKE ?", pattern))
|
||||
}
|
||||
|
||||
// removeEpisodeDerivedContent removes everything produced from an episode analysis.
|
||||
// It intentionally keeps the episode source/subtitle and all project assets.
|
||||
func removeEpisodeDerivedContent(tx *gorm.DB, projectID, episodeID uuid.UUID, media *deletionMediaSet) error {
|
||||
// 封面来自本剧集的首帧,删除反推结果前先清理外键引用,避免随后删除媒体对象时产生约束错误。
|
||||
if err := tx.Exec(`UPDATE creative_projects SET cover_asset_id=NULL
|
||||
WHERE id=? AND cover_asset_id=(SELECT cover_asset_id FROM project_episodes WHERE id=?)`, projectID, episodeID).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Exec("UPDATE project_episodes SET cover_asset_id=NULL WHERE id=?", episodeID).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
taskScope := `(task_type IN ('prompt_reverse','video_generation')
|
||||
OR (task_type='image_generation' AND input_data->>'target_type'='storyboard')) AND
|
||||
(episode_id=? OR storyboard_id IN (SELECT id FROM episode_storyboards WHERE episode_id=?))`
|
||||
if err := media.addQuery(tx.Table("media_assets media").Select("DISTINCT media.id,media.object_key").
|
||||
Joins("JOIN generation_outputs output ON output.media_asset_id=media.id").
|
||||
Joins("JOIN generation_tasks task ON task.id=output.task_id").
|
||||
Where(taskScope, episodeID, episodeID)); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := media.addQuery(tx.Table("media_assets media").Select("DISTINCT media.id,media.object_key").
|
||||
Joins("JOIN episode_storyboards storyboard ON storyboard.thumbnail_asset_id=media.id").
|
||||
Where("storyboard.episode_id=?", episodeID)); err != nil {
|
||||
return err
|
||||
}
|
||||
derivedPrefix := fmt.Sprintf("juyou_ran/video-redraw/projects/%s/episodes/%s/storyboards/%%", projectID, episodeID)
|
||||
if err := collectMediaByObjectKey(tx, media, derivedPrefix); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Exec("UPDATE episode_storyboards SET active_output_id=NULL WHERE episode_id=?", episodeID).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Exec("DELETE FROM generation_outputs WHERE task_id IN (SELECT id FROM generation_tasks WHERE "+taskScope+")", episodeID, episodeID).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Exec("DELETE FROM generation_tasks WHERE "+taskScope, episodeID, episodeID).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Exec(`UPDATE generation_tasks SET storyboard_id=NULL
|
||||
WHERE storyboard_id IN (SELECT id FROM episode_storyboards WHERE episode_id=?)`, episodeID).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.Exec("DELETE FROM episode_storyboards WHERE episode_id=?", episodeID).Error
|
||||
}
|
||||
@@ -0,0 +1,152 @@
|
||||
// 创作媒体服务,封装媒体上传前的所有权检查、关联更新和旧媒体清理事务。
|
||||
package service
|
||||
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"juhe-factory/api/internal/model"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// MediaUploadTarget 汇总上传前需要的业务对象及待替换媒体信息。
|
||||
type MediaUploadTarget struct {
|
||||
Project *model.CreativeProject
|
||||
Episode *model.ProjectEpisode
|
||||
Asset *model.ProjectAsset
|
||||
OldMediaID *uuid.UUID
|
||||
OldObjectKey string
|
||||
}
|
||||
|
||||
// MediaObject 查询当前用户有权播放或下载的媒体对象键。
|
||||
func (s *Creative) MediaObject(userID, mediaID uuid.UUID) (string, error) {
|
||||
var objectKey string
|
||||
err := s.DB.Table("media_assets").
|
||||
Where("id=? AND owner_user_id=? AND deleted_at IS NULL", mediaID, userID).
|
||||
Pluck("object_key", &objectKey).Error
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if strings.TrimSpace(objectKey) == "" {
|
||||
return "", gorm.ErrRecordNotFound
|
||||
}
|
||||
return objectKey, nil
|
||||
}
|
||||
|
||||
// PrepareRedrawMediaUpload 校验重绘项目所有权,并返回指定媒体类型的替换信息。
|
||||
func (s *Creative) PrepareRedrawMediaUpload(userID, projectID uuid.UUID, kind string) (MediaUploadTarget, error) {
|
||||
var project model.CreativeProject
|
||||
if err := s.DB.Where("id=? AND user_id=? AND project_type='video_redraw' AND deleted_at IS NULL", projectID, userID).Take(&project).Error; err != nil {
|
||||
return MediaUploadTarget{}, err
|
||||
}
|
||||
oldMediaID := project.SourceVideoAssetID
|
||||
if kind == "subtitle" {
|
||||
oldMediaID = project.SubtitleAssetID
|
||||
}
|
||||
oldObjectKey, err := s.mediaObjectKey(oldMediaID)
|
||||
if err != nil {
|
||||
return MediaUploadTarget{}, err
|
||||
}
|
||||
return MediaUploadTarget{Project: &project, OldMediaID: oldMediaID, OldObjectKey: oldObjectKey}, nil
|
||||
}
|
||||
|
||||
// PersistRedrawMediaUpload 在事务中保存新媒体、更新项目关联并删除旧媒体记录。
|
||||
func (s *Creative) PersistRedrawMediaUpload(projectID uuid.UUID, kind, currentStatus string, media *model.MediaAsset, oldMediaID *uuid.UUID) error {
|
||||
column, status := "source_video_asset_id", "uploaded"
|
||||
if kind == "subtitle" {
|
||||
column, status = "subtitle_asset_id", currentStatus
|
||||
}
|
||||
return s.DB.Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Create(media).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
updates := map[string]any{column: media.ID, "redraw_status": status}
|
||||
if err := tx.Table("creative_projects").Where("id=?", projectID).Updates(updates).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return deleteReplacedMedia(tx, oldMediaID)
|
||||
})
|
||||
}
|
||||
|
||||
// PrepareEpisodeMediaUpload 校验剧集与项目所有权,并返回指定媒体类型的替换信息。
|
||||
func (s *Creative) PrepareEpisodeMediaUpload(userID, projectID, episodeID uuid.UUID, kind string) (MediaUploadTarget, error) {
|
||||
var episode model.ProjectEpisode
|
||||
err := s.DB.Table("project_episodes e").Joins("JOIN creative_projects p ON p.id=e.project_id").
|
||||
Where("e.id=? AND e.project_id=? AND e.deleted_at IS NULL AND p.user_id=? AND p.deleted_at IS NULL", episodeID, projectID, userID).
|
||||
Select("e.*").First(&episode).Error
|
||||
if err != nil {
|
||||
return MediaUploadTarget{}, err
|
||||
}
|
||||
oldMediaID := episode.SourceVideoAssetID
|
||||
if kind == "subtitle" {
|
||||
oldMediaID = episode.SubtitleAssetID
|
||||
}
|
||||
oldObjectKey, err := s.mediaObjectKey(oldMediaID)
|
||||
if err != nil {
|
||||
return MediaUploadTarget{}, err
|
||||
}
|
||||
return MediaUploadTarget{Episode: &episode, OldMediaID: oldMediaID, OldObjectKey: oldObjectKey}, nil
|
||||
}
|
||||
|
||||
// PersistEpisodeMediaUpload 在事务中保存新媒体、更新剧集关联并删除旧媒体记录。
|
||||
func (s *Creative) PersistEpisodeMediaUpload(episodeID uuid.UUID, kind string, status string, media *model.MediaAsset, oldMediaID *uuid.UUID) error {
|
||||
column := "source_video_asset_id"
|
||||
if kind == "subtitle" {
|
||||
column = "subtitle_asset_id"
|
||||
}
|
||||
return s.DB.Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Create(media).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Table("project_episodes").Where("id=?", episodeID).Updates(map[string]any{column: media.ID, "status": status}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return deleteReplacedMedia(tx, oldMediaID)
|
||||
})
|
||||
}
|
||||
|
||||
// PrepareAssetMediaUpload 校验普通创作项目中的资产所有权并返回资产信息。
|
||||
func (s *Creative) PrepareAssetMediaUpload(userID, projectID, assetID uuid.UUID) (MediaUploadTarget, error) {
|
||||
var asset model.ProjectAsset
|
||||
err := s.DB.Table("project_assets a").Joins("JOIN creative_projects p ON p.id=a.project_id").
|
||||
Where("a.id=? AND a.project_id=? AND a.deleted_at IS NULL AND p.user_id=? AND p.project_type<>'video_redraw' AND p.deleted_at IS NULL", assetID, projectID, userID).
|
||||
Select("a.*").First(&asset).Error
|
||||
if err != nil {
|
||||
return MediaUploadTarget{}, err
|
||||
}
|
||||
return MediaUploadTarget{Asset: &asset}, nil
|
||||
}
|
||||
|
||||
// AttachUploadedAssetAudio 在事务中保存上传音频并更新资产关联。
|
||||
func (s *Creative) AttachUploadedAssetAudio(assetID uuid.UUID, media *model.MediaAsset) error {
|
||||
return s.DB.Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Create(media).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.Table("project_assets").Where("id=?", assetID).Update("audio_asset_id", media.ID).Error
|
||||
})
|
||||
}
|
||||
|
||||
// mediaObjectKey 查询待替换媒体的对象键;空媒体标识表示无需清理。
|
||||
func (s *Creative) mediaObjectKey(mediaID *uuid.UUID) (string, error) {
|
||||
if mediaID == nil {
|
||||
return "", nil
|
||||
}
|
||||
var objectKey string
|
||||
if err := s.DB.Table("media_assets").Where("id=?", *mediaID).Pluck("object_key", &objectKey).Error; err != nil {
|
||||
return "", err
|
||||
}
|
||||
return objectKey, nil
|
||||
}
|
||||
|
||||
// deleteReplacedMedia 删除旧媒体的渠道缓存和媒体记录,供上传事务复用。
|
||||
func deleteReplacedMedia(tx *gorm.DB, mediaID *uuid.UUID) error {
|
||||
if mediaID == nil {
|
||||
return nil
|
||||
}
|
||||
if err := tx.Exec("DELETE FROM channel_asset_cache WHERE media_asset_id=?", *mediaID).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.Exec("DELETE FROM media_assets WHERE id=?", *mediaID).Error
|
||||
}
|
||||
@@ -0,0 +1,73 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"testing"
|
||||
|
||||
"github.com/google/uuid"
|
||||
)
|
||||
|
||||
func TestAssetMentionIndex(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
content string
|
||||
asset string
|
||||
want int
|
||||
}{
|
||||
{name: "space terminated mention", content: "镜头跟随 @丫丫 向前移动", asset: "丫丫", want: 13},
|
||||
{name: "punctuation terminated mention", content: "看向@客厅,停顿。", asset: "客厅", want: 6},
|
||||
{name: "end terminated mention", content: "参考@分镜 1 分镜图", asset: "分镜 1 分镜图", want: 6},
|
||||
{name: "reject shorter prefix", content: "@父亲 走进房间", asset: "父", want: -1},
|
||||
{name: "reject plain asset name", content: "父亲走进房间", asset: "父亲", want: -1},
|
||||
}
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
if got := assetMentionIndex(test.content, test.asset); got != test.want {
|
||||
t.Fatalf("assetMentionIndex(%q, %q) = %d, want %d", test.content, test.asset, got, test.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestRenameAssetMentions(t *testing.T) {
|
||||
tests := []struct {
|
||||
name, content, oldName, newName, want string
|
||||
protectedNames []string
|
||||
}{
|
||||
{name: "all complete mentions", content: "@父亲 走向@客厅,随后看向@父亲", oldName: "父亲", newName: "老周", want: "@老周 走向@客厅,随后看向@老周"},
|
||||
{name: "Chinese description without separator", content: "@林溪微微睁大眼睛", oldName: "林溪", newName: "林溪月", want: "@林溪月微微睁大眼睛"},
|
||||
{name: "reject known longer asset", content: "@父亲 走进房间", oldName: "父", newName: "老周", protectedNames: []string{"父亲"}, want: "@父亲 走进房间"},
|
||||
{name: "ignore plain name", content: "父亲走进房间", oldName: "父亲", newName: "老周", want: "父亲走进房间"},
|
||||
{name: "punctuation in name", content: "参考@A+B。", oldName: "A+B", newName: "组合", want: "参考@组合。"},
|
||||
}
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
if got := renameAssetMentions(test.content, test.oldName, test.newName, test.protectedNames); got != test.want {
|
||||
t.Fatalf("renameAssetMentions() = %q, want %q", got, test.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestRenameAssetRef(t *testing.T) {
|
||||
assetID := uuid.New()
|
||||
otherID := uuid.NewString()
|
||||
value := json.RawMessage(`[{"id":"` + assetID.String() + `","name":"父亲","type":"character"},{"id":"` + otherID + `","name":"客厅"}]`)
|
||||
got, changed, err := renameAssetRef(value, assetID, "老周")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var refs []map[string]any
|
||||
if unmarshalErr := json.Unmarshal(got, &refs); unmarshalErr != nil || !changed || refs[0]["name"] != "老周" || refs[1]["name"] != "客厅" || refs[1]["id"] != otherID {
|
||||
t.Fatalf("renameAssetRef() = %s, changed=%v, err=%v", got, changed, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAssetRefNamesKeepsHistoricalMentionName(t *testing.T) {
|
||||
assetID := uuid.New()
|
||||
value := json.RawMessage(`[{"id":"` + assetID.String() + `","name":"改名前名称","type":"character"}]`)
|
||||
names, err := assetRefNames(value, assetID)
|
||||
if err != nil || len(names) != 1 || names[0] != "改名前名称" {
|
||||
t.Fatalf("assetRefNames() = %#v, err=%v", names, err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,54 @@
|
||||
package service
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestShortDramaProjectDoesNotRequireLocalization(t *testing.T) {
|
||||
input := normalizeCreateProjectInput(ProjectInput{
|
||||
Name: "测试短剧",
|
||||
StyleID: "00000000-0000-0000-0000-000000000001",
|
||||
EraType: "modern_city",
|
||||
AspectRatio: "9:16",
|
||||
})
|
||||
if input.ProjectType != "premium_drama" {
|
||||
t.Fatalf("expected request without localization to be treated as premium_drama, got %q", input.ProjectType)
|
||||
}
|
||||
if err := validateProjectInput(input); err != nil {
|
||||
t.Fatalf("short drama project should not require localization: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestVideoRedrawScriptUsesCreationDefaults(t *testing.T) {
|
||||
followerCount := "百万粉丝"
|
||||
input := normalizeCreateProjectInput(ProjectInput{
|
||||
ProjectType: "video_redraw",
|
||||
Name: "测试剧本",
|
||||
StyleID: "00000000-0000-0000-0000-000000000001",
|
||||
ShortDramaType: "都市情感",
|
||||
PlayCount: "全网播放破亿",
|
||||
AudienceProfile: "25 至 35 岁女性",
|
||||
Producer: "测试出品方",
|
||||
CastMembers: []ProjectCastMemberInput{{
|
||||
Name: "测试主演", FollowerCount: &followerCount,
|
||||
}},
|
||||
})
|
||||
if input.EraType != "modern_city" || input.AspectRatio != "9:16" {
|
||||
t.Fatalf("unexpected redraw creation defaults: era=%q ratio=%q", input.EraType, input.AspectRatio)
|
||||
}
|
||||
if input.Localization == nil || *input.Localization != "china" {
|
||||
t.Fatalf("expected default localization china, got %v", input.Localization)
|
||||
}
|
||||
if err := validateProjectInput(input); err != nil {
|
||||
t.Fatalf("redraw project metadata should be valid after defaults: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestVideoRedrawProjectOnlyRequiresName(t *testing.T) {
|
||||
input := normalizeCreateProjectInput(ProjectInput{
|
||||
ProjectType: "video_redraw",
|
||||
Name: "测试剧本",
|
||||
StyleID: "00000000-0000-0000-0000-000000000001",
|
||||
})
|
||||
if err := validateProjectInput(input); err != nil {
|
||||
t.Fatalf("redraw project optional metadata should be accepted: %v", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,735 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"juhe-factory/api/internal/billing"
|
||||
"juhe-factory/api/internal/model"
|
||||
queuepkg "juhe-factory/api/internal/queue"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/clause"
|
||||
)
|
||||
|
||||
var activeTaskStatuses = []string{"pending_submission", "submitting", "submitted", "processing", "result_ready", "downloading", "cancel_requested"}
|
||||
|
||||
func (s *Creative) InsertStoryboard(userID, projectID, storyboardID uuid.UUID, position string) (*model.EpisodeStoryboard, error) {
|
||||
if position != "before" && position != "after" {
|
||||
return nil, errors.New("插入位置无效")
|
||||
}
|
||||
created := &model.EpisodeStoryboard{}
|
||||
err := s.DB.Transaction(func(tx *gorm.DB) error {
|
||||
var target model.EpisodeStoryboard
|
||||
if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).Table("episode_storyboards sb").Select("sb.*").
|
||||
Joins(`JOIN creative_projects p ON p.id=? AND p.user_id=? AND p.deleted_at IS NULL AND
|
||||
(sb.project_id=p.id OR EXISTS (SELECT 1 FROM project_episodes e WHERE e.id=sb.episode_id AND e.project_id=p.id AND e.deleted_at IS NULL))`, projectID, userID).
|
||||
Where("sb.id=? AND sb.deleted_at IS NULL", storyboardID).Take(&target).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
parentColumn, parentID := storyboardParent(target)
|
||||
sequence, startMS := target.SequenceNo, target.StartMS
|
||||
if position == "after" {
|
||||
sequence, startMS = target.SequenceNo+1, target.EndMS
|
||||
}
|
||||
if err := tx.Exec(`UPDATE episode_storyboards SET sequence_no=sequence_no+100000
|
||||
WHERE `+parentColumn+`=? AND sequence_no>=? AND deleted_at IS NULL`, parentID, sequence).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Exec(`UPDATE episode_storyboards SET sequence_no=sequence_no-99999,start_ms=start_ms+5000,end_ms=end_ms+5000
|
||||
WHERE `+parentColumn+`=? AND sequence_no>=? AND deleted_at IS NULL`, parentID, sequence+100000).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
created = &model.EpisodeStoryboard{
|
||||
EpisodeID: target.EpisodeID, ProjectID: target.ProjectID, SequenceNo: sequence, StableKey: "manual-" + uuid.NewString(),
|
||||
StartMS: startMS, EndMS: startMS + 5000, DurationSeconds: 5, Title: fmt.Sprintf("分镜 %d", sequence),
|
||||
Dialogue: json.RawMessage(`[]`), AssetRefs: json.RawMessage(`[]`), Status: "idle",
|
||||
}
|
||||
return tx.Create(created).Error
|
||||
})
|
||||
return created, err
|
||||
}
|
||||
|
||||
func (s *Creative) DeleteStoryboard(userID, projectID, storyboardID uuid.UUID) ([]string, error) {
|
||||
media := newDeletionMediaSet()
|
||||
err := s.DB.Transaction(func(tx *gorm.DB) error {
|
||||
var storyboard model.EpisodeStoryboard
|
||||
if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).Table("episode_storyboards sb").Select("sb.*").
|
||||
Joins(`JOIN creative_projects p ON p.id=? AND p.user_id=? AND p.deleted_at IS NULL AND
|
||||
(sb.project_id=p.id OR EXISTS (SELECT 1 FROM project_episodes e WHERE e.id=sb.episode_id AND e.project_id=p.id AND e.deleted_at IS NULL))`, projectID, userID).
|
||||
Where("sb.id=? AND sb.deleted_at IS NULL", storyboardID).Take(&storyboard).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
parentColumn, parentID := storyboardParent(storyboard)
|
||||
var active int64
|
||||
if err := tx.Model(&model.GenerationTask{}).Where("storyboard_id=? AND status IN ?", storyboardID, activeTaskStatuses).Count(&active).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if active > 0 {
|
||||
return errors.New("当前分镜存在进行中的任务,暂时不能删除")
|
||||
}
|
||||
if err := media.addQuery(tx.Table("media_assets media").Select("DISTINCT media.id,media.object_key").
|
||||
Joins("JOIN generation_outputs output ON output.media_asset_id=media.id").
|
||||
Joins("JOIN generation_tasks gt ON gt.id=output.task_id").
|
||||
Where("gt.storyboard_id=? AND (gt.task_type IN ('video_generation','prompt_reverse') OR (gt.task_type='image_generation' AND gt.input_data->>'target_type'='storyboard'))", storyboardID)); err != nil {
|
||||
return err
|
||||
}
|
||||
if storyboard.ThumbnailAssetID != nil {
|
||||
if err := media.addQuery(tx.Table("media_assets").Select("id,object_key").Where("id=?", *storyboard.ThumbnailAssetID)); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
storyboardMediaPattern := fmt.Sprintf("juyou_ran/video-redraw/projects/%s/storyboards/%%-%s/%%", projectID, storyboardID)
|
||||
if storyboard.EpisodeID != nil {
|
||||
storyboardMediaPattern = fmt.Sprintf("juyou_ran/video-redraw/projects/%s/episodes/%s/storyboards/%%-%s/%%", projectID, *storyboard.EpisodeID, storyboardID)
|
||||
}
|
||||
if err := collectMediaByObjectKey(tx, media, storyboardMediaPattern); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Model(&storyboard).Update("active_output_id", nil).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Exec(`DELETE FROM generation_outputs WHERE task_id IN (
|
||||
SELECT id FROM generation_tasks WHERE storyboard_id=? AND (task_type IN ('video_generation','prompt_reverse')
|
||||
OR (task_type='image_generation' AND input_data->>'target_type'='storyboard')))`, storyboardID).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Exec(`DELETE FROM generation_tasks WHERE storyboard_id=? AND (task_type IN ('video_generation','prompt_reverse')
|
||||
OR (task_type='image_generation' AND input_data->>'target_type'='storyboard'))`, storyboardID).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Model(&model.GenerationTask{}).Where("storyboard_id=?", storyboardID).Update("storyboard_id", nil).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Exec("DELETE FROM episode_storyboards WHERE id=?", storyboardID).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Exec(`UPDATE episode_storyboards SET sequence_no=sequence_no+100000
|
||||
WHERE `+parentColumn+`=? AND sequence_no>? AND deleted_at IS NULL`, parentID, storyboard.SequenceNo).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
shiftMS := int64(storyboard.DurationSeconds * 1000)
|
||||
if err := tx.Exec(`UPDATE episode_storyboards SET sequence_no=sequence_no-100001,
|
||||
start_ms=GREATEST(0,start_ms-?),end_ms=GREATEST(1,end_ms-?)
|
||||
WHERE `+parentColumn+`=? AND sequence_no>? AND deleted_at IS NULL`, shiftMS, shiftMS, parentID, storyboard.SequenceNo+100000).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return media.deleteRows(tx)
|
||||
})
|
||||
return media.objectKeys(), err
|
||||
}
|
||||
|
||||
func storyboardParent(storyboard model.EpisodeStoryboard) (string, uuid.UUID) {
|
||||
if storyboard.ProjectID != nil {
|
||||
return "project_id", *storyboard.ProjectID
|
||||
}
|
||||
return "episode_id", *storyboard.EpisodeID
|
||||
}
|
||||
|
||||
func (s *Creative) DeleteEpisode(userID, projectID, episodeID uuid.UUID) ([]string, error) {
|
||||
media := newDeletionMediaSet()
|
||||
err := s.DB.Transaction(func(tx *gorm.DB) error {
|
||||
var episode model.ProjectEpisode
|
||||
if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).Table("project_episodes e").Select("e.*").
|
||||
Joins("JOIN creative_projects p ON p.id=e.project_id AND p.user_id=? AND p.deleted_at IS NULL", userID).
|
||||
Where("e.id=? AND e.project_id=? AND e.deleted_at IS NULL", episodeID, projectID).Take(&episode).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
var active int64
|
||||
if err := tx.Model(&model.GenerationTask{}).Where("episode_id=? AND status IN ?", episodeID, activeTaskStatuses).Count(&active).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if active > 0 {
|
||||
return errors.New("剧集存在排队中或处理中的任务,暂时不能删除")
|
||||
}
|
||||
if err := tx.Model(&model.DramaParseTask{}).Where("episode_id=? AND status IN ?", episodeID, []string{"queued", "running", "retry_wait", "cancel_requested"}).Count(&active).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if active > 0 {
|
||||
return errors.New("剧集存在进行中的剧本解析任务,暂时不能删除")
|
||||
}
|
||||
if err := removeEpisodeDerivedContent(tx, projectID, episodeID, media); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := media.addQuery(tx.Table("media_assets media").Select("DISTINCT media.id,media.object_key").
|
||||
Joins(`JOIN project_episodes episode ON media.id=episode.cover_asset_id
|
||||
OR media.id=episode.source_video_asset_id OR media.id=episode.subtitle_asset_id`).
|
||||
Where("episode.id=?", episodeID)); err != nil {
|
||||
return err
|
||||
}
|
||||
// 如果删除的是第一集,将项目封面切换到新的第一集封面。
|
||||
if episode.CoverAssetID != nil {
|
||||
var nextEpisode struct{ CoverAssetID *uuid.UUID }
|
||||
if err := tx.Table("project_episodes").Select("cover_asset_id").Where("project_id=? AND id<>? AND deleted_at IS NULL AND cover_asset_id IS NOT NULL", projectID, episodeID).
|
||||
Order("episode_no").Limit(1).Find(&nextEpisode).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Model(&model.CreativeProject{}).Where("id=?", projectID).Update("cover_asset_id", nextEpisode.CoverAssetID).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
episodeMediaPattern := fmt.Sprintf("juyou_ran/video-redraw/projects/%s/episodes/%%%s/%%", projectID, episodeID)
|
||||
if err := collectMediaByObjectKey(tx, media, episodeMediaPattern); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Exec("UPDATE generation_tasks SET episode_id=NULL WHERE episode_id=?", episodeID).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Exec("UPDATE project_assets SET source_episode_id=NULL WHERE source_episode_id=?", episodeID).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Exec("UPDATE drama_parse_batches SET current_episode_id=NULL WHERE current_episode_id=?", episodeID).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Exec("DELETE FROM project_episodes WHERE id=?", episodeID).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return media.deleteRows(tx)
|
||||
})
|
||||
return media.objectKeys(), err
|
||||
}
|
||||
|
||||
// ResetEpisodeAnalysis permanently removes generated storyboards, prompts and their task history.
|
||||
// Project assets are intentionally left untouched so they can be reused for a new analysis.
|
||||
func (s *Creative) ResetEpisodeAnalysis(userID, projectID, episodeID uuid.UUID) ([]string, error) {
|
||||
media := newDeletionMediaSet()
|
||||
err := s.DB.Transaction(func(tx *gorm.DB) error {
|
||||
var episode model.ProjectEpisode
|
||||
if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).Table("project_episodes e").Select("e.*").
|
||||
Joins("JOIN creative_projects p ON p.id=e.project_id AND p.user_id=? AND p.deleted_at IS NULL", userID).
|
||||
Where("e.id=? AND e.project_id=? AND e.deleted_at IS NULL", episodeID, projectID).Take(&episode).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
var active int64
|
||||
if err := tx.Model(&model.GenerationTask{}).Where("episode_id=? AND status IN ?", episodeID, activeTaskStatuses).Count(&active).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if active > 0 {
|
||||
return errors.New("当前剧集存在进行中的生成任务,暂时不能重新解析")
|
||||
}
|
||||
if err := tx.Model(&model.DramaParseTask{}).Where("episode_id=? AND status IN ?", episodeID, []string{"queued", "running", "retry_wait", "cancel_requested"}).Count(&active).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if active > 0 {
|
||||
return errors.New("当前剧集存在进行中的剧本解析任务,暂时不能重新解析")
|
||||
}
|
||||
if err := removeEpisodeDerivedContent(tx, projectID, episodeID, media); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := media.deleteRows(tx); err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.Model(&episode).Updates(map[string]any{"status": "uploaded", "analysis_message": nil, "redraw_script": nil}).Error
|
||||
})
|
||||
return media.objectKeys(), err
|
||||
}
|
||||
|
||||
func (s *Creative) ListTasks(userID, projectID uuid.UUID, episodeID *uuid.UUID) ([]map[string]any, error) {
|
||||
items := make([]map[string]any, 0)
|
||||
query := s.DB.Table("generation_tasks gt").
|
||||
Select(`gt.id,gt.request_id,gt.task_type,gt.status,gt.episode_id,gt.storyboard_id,gt.input_data->>'mode' AS mode,gt.input_data->>'phase_started_at' AS phase_started_at,gt.estimated_points::text AS estimated_points,
|
||||
gt.actual_points::text AS actual_points,gt.error_code,gt.error_message,gt.created_at,gt.submitted_at,gt.finished_at,
|
||||
coalesce(sb.sequence_no,(gt.input_data->>'current_sequence')::int,0) AS sequence_no,gt.input_data->>'asset_id' AS asset_id,media.public_url AS result_url,
|
||||
(SELECT max(progress.created_at) FROM generation_outputs progress WHERE progress.task_id=gt.id) AS last_result_at`).
|
||||
Joins("JOIN creative_projects p ON p.id=gt.project_id AND p.user_id=? AND p.deleted_at IS NULL", userID).
|
||||
Joins("LEFT JOIN episode_storyboards sb ON sb.id=gt.storyboard_id").
|
||||
Joins("LEFT JOIN generation_outputs output ON output.task_id=gt.id AND output.sequence_no=1").
|
||||
Joins("LEFT JOIN media_assets media ON media.id=output.media_asset_id AND media.deleted_at IS NULL").
|
||||
Where("gt.project_id=?", projectID)
|
||||
if episodeID != nil {
|
||||
query = query.Where("gt.episode_id=?", *episodeID)
|
||||
}
|
||||
return items, query.Order("gt.created_at DESC").Limit(500).Find(&items).Error
|
||||
}
|
||||
|
||||
func (s *Creative) CancelTask(userID, projectID, taskID uuid.UUID) error {
|
||||
var channelID *uuid.UUID
|
||||
err := s.DB.Transaction(func(tx *gorm.DB) error {
|
||||
var task model.GenerationTask
|
||||
if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).Where("id=? AND project_id=? AND user_id=?", taskID, projectID, userID).First(&task).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
var projectType string
|
||||
if err := tx.Table("creative_projects").Where("id=?", projectID).Pluck("project_type", &projectType).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
channelID = task.ChannelID
|
||||
if task.TaskType == "prompt_reverse" && (task.Status == "pending_submission" || task.Status == "submitted" || task.Status == "processing" || task.Status == "cancel_requested") {
|
||||
return cancelPromptReverseTask(tx, &task)
|
||||
}
|
||||
switch task.Status {
|
||||
case "pending_submission":
|
||||
refunded := false
|
||||
chargeCancellation := task.TaskType == "video_generation" || projectType == "premium_drama"
|
||||
if !chargeCancellation {
|
||||
var err error
|
||||
refunded, err = billing.RefundGenerationTask(tx, &task, "视频生成失败返还")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
actual := "0.00"
|
||||
if chargeCancellation {
|
||||
actual = task.PrepaidPoints
|
||||
}
|
||||
updates := map[string]any{"status": "cancelled", "actual_points": actual, "finished_at": time.Now(), "error_code": nil, "error_message": "用户取消"}
|
||||
if refunded {
|
||||
updates["cost_refunded"] = true
|
||||
}
|
||||
if err := tx.Model(&task).Updates(updates).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return resetRelatedStatus(tx, task)
|
||||
case "submitted", "processing":
|
||||
if task.TaskType != "prompt_reverse" && strings.TrimSpace(task.UpstreamTaskID) == "" {
|
||||
return errors.New("上游任务标识尚未就绪,请稍后重试取消")
|
||||
}
|
||||
return tx.Model(&task).Updates(map[string]any{"status": "cancel_requested", "error_code": nil, "error_message": "等待当前处理安全结束"}).Error
|
||||
case "submitting":
|
||||
return errors.New("任务正在提交,请稍后重试取消")
|
||||
case "cancel_requested":
|
||||
return nil
|
||||
case "result_ready", "downloading", "succeeded":
|
||||
return errors.New("生成已经完成,不能取消")
|
||||
case "failed", "cancelled":
|
||||
return nil
|
||||
default:
|
||||
return fmt.Errorf("任务状态 %s 不允许取消", task.Status)
|
||||
}
|
||||
})
|
||||
if err == nil && channelID != nil && s.Queue != nil {
|
||||
_ = queuepkg.EnqueueID(s.Queue, queuepkg.TypeDispatchChannel, *channelID, 0)
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
func cancelPromptReverseTask(tx *gorm.DB, task *model.GenerationTask) error {
|
||||
var input struct {
|
||||
Mode string `json:"mode"`
|
||||
}
|
||||
_ = json.Unmarshal(task.InputData, &input)
|
||||
updates := map[string]any{
|
||||
"status": "cancelled", "actual_points": task.PrepaidPoints, "finished_at": time.Now(),
|
||||
"error_code": nil, "error_message": "用户取消,已扣积分不退",
|
||||
}
|
||||
if err := tx.Model(task).Updates(updates).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if task.StoryboardID != nil {
|
||||
return nil
|
||||
}
|
||||
if task.EpisodeID == nil && task.ProjectID != nil {
|
||||
if input.Mode == "script" {
|
||||
return tx.Model(&model.CreativeProject{}).Where("id=?", *task.ProjectID).
|
||||
Updates(map[string]any{"redraw_status": "review", "analysis_message": "剧本反推已取消"}).Error
|
||||
}
|
||||
var storyboardCount int64
|
||||
if err := tx.Model(&model.EpisodeStoryboard{}).Where("project_id=? AND deleted_at IS NULL", *task.ProjectID).Count(&storyboardCount).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
status := "uploaded"
|
||||
if storyboardCount > 0 {
|
||||
status = "review"
|
||||
}
|
||||
return tx.Model(&model.CreativeProject{}).Where("id=?", *task.ProjectID).Updates(map[string]any{"redraw_status": status, "analysis_message": "视频分析已取消"}).Error
|
||||
}
|
||||
if task.EpisodeID == nil {
|
||||
return nil
|
||||
}
|
||||
var storyboardCount int64
|
||||
if err := tx.Model(&model.EpisodeStoryboard{}).Where("episode_id=? AND deleted_at IS NULL", *task.EpisodeID).Count(&storyboardCount).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
status := "uploaded"
|
||||
if storyboardCount > 0 {
|
||||
status = "review"
|
||||
}
|
||||
return tx.Model(&model.ProjectEpisode{}).Where("id=?", *task.EpisodeID).Updates(map[string]any{"status": status, "analysis_message": "视频分析已取消"}).Error
|
||||
}
|
||||
|
||||
func (s *Creative) failQueuedTask(taskID uuid.UUID, code, message string) error {
|
||||
return s.failQueuedTasks([]uuid.UUID{taskID}, code, message)
|
||||
}
|
||||
|
||||
func (s *Creative) failQueuedTasks(taskIDs []uuid.UUID, code, message string) error {
|
||||
return s.DB.Transaction(func(tx *gorm.DB) error {
|
||||
for _, taskID := range taskIDs {
|
||||
var task model.GenerationTask
|
||||
if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).Where("id=?", taskID).First(&task).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if task.Status == "succeeded" || task.Status == "failed" || task.Status == "cancelled" {
|
||||
continue
|
||||
}
|
||||
refundRemark := map[string]string{"prompt_reverse": "视频反推失败返还", "script_analysis": "剧本分析失败返还", "image_generation": "图片生成失败返还", "video_generation": "视频生成失败返还"}[task.TaskType]
|
||||
var refunded bool
|
||||
var err error
|
||||
if task.TaskType == "prompt_reverse" || task.TaskType == "script_analysis" {
|
||||
refunded, err = billing.RefundTextGenerationTask(tx, &task, refundRemark)
|
||||
} else {
|
||||
refunded, err = billing.RefundGenerationTask(tx, &task, refundRemark)
|
||||
}
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
updates := map[string]any{
|
||||
"status": "failed", "actual_points": "0.00", "finished_at": time.Now(),
|
||||
"error_code": code, "error_message": message,
|
||||
}
|
||||
if refunded {
|
||||
updates["cost_refunded"] = true
|
||||
}
|
||||
if err := tx.Model(&task).Updates(updates).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if task.TaskType == "video_generation" && task.StoryboardID != nil {
|
||||
if err := tx.Model(&model.EpisodeStoryboard{}).Where("id=?", *task.StoryboardID).Update("status", "idle").Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
if task.TaskType == "prompt_reverse" && task.StoryboardID == nil && task.EpisodeID != nil {
|
||||
var storyboardCount int64
|
||||
if err := tx.Model(&model.EpisodeStoryboard{}).Where("episode_id=? AND deleted_at IS NULL", *task.EpisodeID).Count(&storyboardCount).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
status := "uploaded"
|
||||
if storyboardCount > 0 {
|
||||
status = "review"
|
||||
}
|
||||
if err := tx.Model(&model.ProjectEpisode{}).Where("id=?", *task.EpisodeID).Updates(map[string]any{"status": status, "analysis_message": message}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
} else if task.TaskType == "prompt_reverse" && task.StoryboardID == nil && task.ProjectID != nil {
|
||||
var storyboardCount int64
|
||||
if err := tx.Model(&model.EpisodeStoryboard{}).Where("project_id=? AND deleted_at IS NULL", *task.ProjectID).Count(&storyboardCount).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
status := "uploaded"
|
||||
if storyboardCount > 0 {
|
||||
status = "review"
|
||||
}
|
||||
if err := tx.Model(&model.CreativeProject{}).Where("id=?", *task.ProjectID).Updates(map[string]any{"redraw_status": status, "analysis_message": message}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
})
|
||||
}
|
||||
|
||||
func resetRelatedStatus(tx *gorm.DB, task model.GenerationTask) error {
|
||||
if task.TaskType == "video_generation" && task.StoryboardID != nil {
|
||||
var active int64
|
||||
if err := tx.Model(&model.GenerationTask{}).Where("storyboard_id=? AND task_type='video_generation' AND status IN ?", *task.StoryboardID, activeTaskStatuses).Count(&active).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if active == 0 {
|
||||
return tx.Model(&model.EpisodeStoryboard{}).Where("id=?", *task.StoryboardID).Update("status", "idle").Error
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *Creative) ListStoryboardOutputs(userID, projectID, storyboardID uuid.UUID) ([]map[string]any, error) {
|
||||
items := make([]map[string]any, 0)
|
||||
err := s.DB.Table("generation_outputs output").
|
||||
Select(`output.id,output.media_asset_id,output.output_type,output.sequence_no,output.metadata,output.created_at,output.task_id,
|
||||
media.public_url,media.mime_type,media.size_bytes,gt.estimated_points::text AS estimated_points,
|
||||
gt.actual_points::text AS actual_points,CASE WHEN gt.task_type='video_generation' THEN sb.active_output_id=output.id ELSE sb.thumbnail_asset_id=output.media_asset_id END AS active,
|
||||
CASE WHEN gt.task_type='video_generation' THEN sb.active_output_id=output.id OR coalesce(output.metadata->>'candidate','false')='true' ELSE false END AS candidate`).
|
||||
Joins("JOIN generation_tasks gt ON gt.id=output.task_id AND gt.storyboard_id=? AND (gt.task_type='video_generation' OR (gt.task_type='image_generation' AND gt.input_data->>'target_type'='storyboard'))", storyboardID).
|
||||
Joins("JOIN media_assets media ON media.id=output.media_asset_id AND media.deleted_at IS NULL").
|
||||
Joins("JOIN episode_storyboards sb ON sb.id=gt.storyboard_id AND sb.deleted_at IS NULL").
|
||||
Joins(`JOIN creative_projects p ON p.id=? AND p.user_id=? AND p.deleted_at IS NULL AND
|
||||
(sb.project_id=p.id OR EXISTS (SELECT 1 FROM project_episodes e WHERE e.id=sb.episode_id AND e.project_id=p.id AND e.deleted_at IS NULL))`, projectID, userID).
|
||||
Order("output.created_at DESC").Find(&items).Error
|
||||
return items, err
|
||||
}
|
||||
|
||||
// RemoveStoryboardImage 清空分镜主图引用,保留生成结果以便在历史记录中恢复。
|
||||
func (s *Creative) RemoveStoryboardImage(userID, projectID, storyboardID uuid.UUID) error {
|
||||
ownedStoryboard := `(project_id=? AND EXISTS (SELECT 1 FROM creative_projects p WHERE p.id=? AND p.user_id=? AND p.deleted_at IS NULL)) OR
|
||||
(episode_id IN (SELECT e.id FROM project_episodes e JOIN creative_projects p ON p.id=e.project_id
|
||||
WHERE e.project_id=? AND p.user_id=? AND e.deleted_at IS NULL AND p.deleted_at IS NULL))`
|
||||
result := s.DB.Model(&model.EpisodeStoryboard{}).
|
||||
Where("id=? AND thumbnail_asset_id IS NOT NULL AND deleted_at IS NULL AND ("+ownedStoryboard+")", storyboardID, projectID, projectID, userID, projectID, userID).
|
||||
Updates(map[string]any{"thumbnail_asset_id": nil, "status": "idle", "user_edited": true})
|
||||
if result.Error != nil {
|
||||
return result.Error
|
||||
}
|
||||
if result.RowsAffected == 0 {
|
||||
var exists int64
|
||||
if err := s.DB.Model(&model.EpisodeStoryboard{}).
|
||||
Where("id=? AND deleted_at IS NULL AND ("+ownedStoryboard+")", storyboardID, projectID, projectID, userID, projectID, userID).
|
||||
Count(&exists).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if exists == 0 {
|
||||
return gorm.ErrRecordNotFound
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *Creative) AddStoryboardCandidate(userID, projectID, storyboardID, outputID uuid.UUID) error {
|
||||
return s.DB.Transaction(func(tx *gorm.DB) error {
|
||||
var row struct{ Candidate bool }
|
||||
if err := tx.Table("generation_outputs output").
|
||||
Select("sb.active_output_id=output.id OR coalesce(output.metadata->>'candidate','false')='true' AS candidate").
|
||||
Joins("JOIN generation_tasks gt ON gt.id=output.task_id AND gt.storyboard_id=? AND gt.task_type='video_generation'", storyboardID).
|
||||
Joins("JOIN episode_storyboards sb ON sb.id=gt.storyboard_id AND sb.deleted_at IS NULL").
|
||||
Joins("JOIN project_episodes e ON e.id=sb.episode_id AND e.deleted_at IS NULL").
|
||||
Joins("JOIN creative_projects p ON p.id=e.project_id AND p.id=? AND p.user_id=? AND p.project_type='premium_drama' AND p.deleted_at IS NULL", projectID, userID).
|
||||
Where("output.id=?", outputID).Take(&row).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if row.Candidate {
|
||||
return nil
|
||||
}
|
||||
var count int64
|
||||
if err := tx.Table("generation_outputs output").
|
||||
Joins("JOIN generation_tasks gt ON gt.id=output.task_id AND gt.storyboard_id=? AND gt.task_type='video_generation'", storyboardID).
|
||||
Joins("JOIN episode_storyboards sb ON sb.id=gt.storyboard_id").
|
||||
Where("sb.active_output_id=output.id OR coalesce(output.metadata->>'candidate','false')='true'").Count(&count).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if count >= 3 {
|
||||
return errors.New("当前分镜最多保留 3 个备选视频,请先删除一个备选视频")
|
||||
}
|
||||
var activeTasks int64
|
||||
if err := tx.Model(&model.GenerationTask{}).
|
||||
Where("storyboard_id=? AND task_type='video_generation' AND status IN ?", storyboardID, activeTaskStatuses).
|
||||
Count(&activeTasks).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if count+activeTasks >= 3 {
|
||||
return errors.New("生成中的视频已占用备选位置,请等待生成完成或先删除一个备选视频")
|
||||
}
|
||||
return tx.Model(&model.GenerationOutput{}).Where("id=?", outputID).
|
||||
Update("metadata", gorm.Expr("jsonb_set(coalesce(metadata,'{}'::jsonb),'{candidate}','true'::jsonb,true)")).Error
|
||||
})
|
||||
}
|
||||
|
||||
func (s *Creative) RemoveStoryboardCandidate(userID, projectID, storyboardID, outputID uuid.UUID) error {
|
||||
return s.DB.Transaction(func(tx *gorm.DB) error {
|
||||
var row struct {
|
||||
Active bool
|
||||
Candidate bool
|
||||
}
|
||||
if err := tx.Table("generation_outputs output").
|
||||
Select("sb.active_output_id=output.id AS active,sb.active_output_id=output.id OR coalesce(output.metadata->>'candidate','false')='true' AS candidate").
|
||||
Joins("JOIN generation_tasks gt ON gt.id=output.task_id AND gt.storyboard_id=? AND gt.task_type='video_generation'", storyboardID).
|
||||
Joins("JOIN episode_storyboards sb ON sb.id=gt.storyboard_id AND sb.deleted_at IS NULL").
|
||||
Joins("JOIN project_episodes e ON e.id=sb.episode_id AND e.deleted_at IS NULL").
|
||||
Joins("JOIN creative_projects p ON p.id=e.project_id AND p.id=? AND p.user_id=? AND p.project_type='premium_drama' AND p.deleted_at IS NULL", projectID, userID).
|
||||
Where("output.id=?", outputID).Take(&row).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if !row.Candidate {
|
||||
return errors.New("该视频不在备选列表中")
|
||||
}
|
||||
if row.Active {
|
||||
var replacement struct{ ID uuid.UUID }
|
||||
if err := tx.Table("generation_outputs output").Select("output.id").
|
||||
Joins("JOIN generation_tasks gt ON gt.id=output.task_id AND gt.storyboard_id=? AND gt.task_type='video_generation'", storyboardID).
|
||||
Where("output.id<>? AND coalesce(output.metadata->>'candidate','false')='true'", outputID).
|
||||
Order("output.created_at DESC").Limit(1).Scan(&replacement).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
var replacementID *uuid.UUID
|
||||
if replacement.ID != uuid.Nil {
|
||||
replacementID = &replacement.ID
|
||||
}
|
||||
status := "idle"
|
||||
if replacementID != nil {
|
||||
status = "completed"
|
||||
}
|
||||
if err := tx.Model(&model.EpisodeStoryboard{}).Where("id=?", storyboardID).
|
||||
Updates(map[string]any{"active_output_id": replacementID, "status": status}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return tx.Model(&model.GenerationOutput{}).Where("id=?", outputID).
|
||||
Update("metadata", gorm.Expr("coalesce(metadata,'{}'::jsonb)-'candidate'")).Error
|
||||
})
|
||||
}
|
||||
|
||||
func (s *Creative) ActivateStoryboardOutput(userID, projectID, storyboardID, outputID uuid.UUID) error {
|
||||
var output struct {
|
||||
MediaAssetID uuid.UUID
|
||||
TaskType string
|
||||
}
|
||||
if err := s.DB.Table("generation_outputs output").Select("output.media_asset_id,gt.task_type").
|
||||
Joins("JOIN generation_tasks gt ON gt.id=output.task_id AND gt.storyboard_id=? AND (gt.task_type='video_generation' OR (gt.task_type='image_generation' AND gt.input_data->>'target_type'='storyboard'))", storyboardID).
|
||||
Joins("JOIN episode_storyboards sb ON sb.id=gt.storyboard_id AND sb.deleted_at IS NULL").
|
||||
Joins(`JOIN creative_projects p ON p.id=? AND p.user_id=? AND p.deleted_at IS NULL AND
|
||||
(sb.project_id=p.id OR EXISTS (SELECT 1 FROM project_episodes e WHERE e.id=sb.episode_id AND e.project_id=p.id AND e.deleted_at IS NULL))`, projectID, userID).
|
||||
Where("output.id=?", outputID).Take(&output).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
updates := map[string]any{"status": "completed", "updated_at": time.Now()}
|
||||
if output.TaskType == "video_generation" {
|
||||
updates["active_output_id"] = outputID
|
||||
} else {
|
||||
updates["thumbnail_asset_id"] = output.MediaAssetID
|
||||
}
|
||||
result := s.DB.Model(&model.EpisodeStoryboard{}).Where("id=? AND deleted_at IS NULL", storyboardID).Updates(updates)
|
||||
if result.Error != nil {
|
||||
return result.Error
|
||||
}
|
||||
if result.RowsAffected == 0 {
|
||||
return gorm.ErrRecordNotFound
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *Creative) DeleteStoryboardOutput(userID, projectID, storyboardID, outputID uuid.UUID) (string, error) {
|
||||
var objectKey string
|
||||
err := s.DB.Transaction(func(tx *gorm.DB) error {
|
||||
var row struct {
|
||||
TaskID uuid.UUID
|
||||
MediaID uuid.UUID
|
||||
ObjectKey string
|
||||
Status string
|
||||
Candidate bool
|
||||
}
|
||||
if err := tx.Table("generation_outputs output").Select("output.task_id,output.media_asset_id AS media_id,media.object_key,gt.status,sb.active_output_id=output.id OR coalesce(output.metadata->>'candidate','false')='true' AS candidate").
|
||||
Joins("JOIN generation_tasks gt ON gt.id=output.task_id AND gt.storyboard_id=? AND (gt.task_type='video_generation' OR (gt.task_type='image_generation' AND gt.input_data->>'target_type'='storyboard'))", storyboardID).
|
||||
Joins("JOIN media_assets media ON media.id=output.media_asset_id").
|
||||
Joins("JOIN episode_storyboards sb ON sb.id=gt.storyboard_id").
|
||||
Joins(`JOIN creative_projects p ON p.id=? AND p.user_id=? AND p.deleted_at IS NULL AND
|
||||
(sb.project_id=p.id OR EXISTS (SELECT 1 FROM project_episodes e WHERE e.id=sb.episode_id AND e.project_id=p.id AND e.deleted_at IS NULL))`, projectID, userID).
|
||||
Where("output.id=?", outputID).Take(&row).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if row.Status == "pending_submission" || row.Status == "submitting" || row.Status == "submitted" || row.Status == "processing" || row.Status == "cancel_requested" {
|
||||
return errors.New("视频仍在生成中,暂时不能删除历史记录")
|
||||
}
|
||||
if row.Candidate {
|
||||
return errors.New("备选视频不能直接删除,请先移入历史记录")
|
||||
}
|
||||
if err := tx.Exec("UPDATE episode_storyboards SET active_output_id=NULL,status='idle' WHERE id=? AND active_output_id=?", storyboardID, outputID).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Exec("UPDATE episode_storyboards SET thumbnail_asset_id=NULL,status='idle' WHERE id=? AND thumbnail_asset_id=?", storyboardID, row.MediaID).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Exec("DELETE FROM generation_outputs WHERE id=?", outputID).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Exec("DELETE FROM generation_tasks WHERE id=?", row.TaskID).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Exec("DELETE FROM channel_asset_cache WHERE media_asset_id=?", row.MediaID).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Exec("DELETE FROM media_assets WHERE id=?", row.MediaID).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
objectKey = row.ObjectKey
|
||||
return nil
|
||||
})
|
||||
return objectKey, err
|
||||
}
|
||||
|
||||
func (s *Creative) ListAssetOutputs(userID, projectID, assetID uuid.UUID) ([]map[string]any, error) {
|
||||
items := make([]map[string]any, 0)
|
||||
err := s.DB.Table("generation_outputs output").
|
||||
Select(`output.id,output.media_asset_id,output.output_type,output.sequence_no,output.metadata,output.created_at,output.task_id,
|
||||
media.public_url,media.mime_type,media.size_bytes,gt.estimated_points::text AS estimated_points,
|
||||
gt.actual_points::text AS actual_points,(a.image_asset_id=output.media_asset_id) AS active`).
|
||||
Joins("JOIN generation_tasks gt ON gt.id=output.task_id AND gt.task_type='image_generation' AND gt.input_data->>'asset_id'=?", assetID.String()).
|
||||
Joins("JOIN media_assets media ON media.id=output.media_asset_id AND media.deleted_at IS NULL").
|
||||
Joins("JOIN project_assets a ON a.id=? AND a.project_id=? AND a.deleted_at IS NULL", assetID, projectID).
|
||||
Joins("JOIN creative_projects p ON p.id=a.project_id AND p.user_id=? AND p.deleted_at IS NULL", userID).
|
||||
Order("output.created_at DESC").Find(&items).Error
|
||||
return items, err
|
||||
}
|
||||
|
||||
func (s *Creative) ActivateAssetOutput(userID, projectID, assetID, outputID uuid.UUID) error {
|
||||
result := s.DB.Exec(`UPDATE project_assets a SET image_asset_id=(SELECT media_asset_id FROM generation_outputs WHERE id=?),updated_at=CURRENT_TIMESTAMP
|
||||
WHERE a.id=? AND a.project_id=? AND a.deleted_at IS NULL AND EXISTS (
|
||||
SELECT 1 FROM generation_outputs output JOIN generation_tasks gt ON gt.id=output.task_id
|
||||
JOIN creative_projects p ON p.id=a.project_id
|
||||
WHERE output.id=? AND gt.task_type='image_generation' AND gt.input_data->>'asset_id'=? AND p.user_id=? AND p.deleted_at IS NULL)`,
|
||||
outputID, assetID, projectID, outputID, assetID.String(), userID)
|
||||
if result.Error != nil {
|
||||
return result.Error
|
||||
}
|
||||
if result.RowsAffected == 0 {
|
||||
return gorm.ErrRecordNotFound
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *Creative) DeleteAssetOutput(userID, projectID, assetID, outputID uuid.UUID) (string, error) {
|
||||
var objectKey string
|
||||
err := s.DB.Transaction(func(tx *gorm.DB) error {
|
||||
var row struct {
|
||||
TaskID uuid.UUID
|
||||
MediaID uuid.UUID
|
||||
ObjectKey string
|
||||
Status string
|
||||
}
|
||||
if err := tx.Table("generation_outputs output").Select("output.task_id,output.media_asset_id AS media_id,media.object_key,gt.status").
|
||||
Joins("JOIN generation_tasks gt ON gt.id=output.task_id AND gt.task_type='image_generation' AND gt.input_data->>'asset_id'=?", assetID.String()).
|
||||
Joins("JOIN media_assets media ON media.id=output.media_asset_id").
|
||||
Joins("JOIN project_assets a ON a.id=? AND a.project_id=? AND a.deleted_at IS NULL", assetID, projectID).
|
||||
Joins("JOIN creative_projects p ON p.id=a.project_id AND p.user_id=? AND p.deleted_at IS NULL", userID).
|
||||
Where("output.id=?", outputID).Take(&row).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
for _, status := range activeTaskStatuses {
|
||||
if status == row.Status {
|
||||
return errors.New("鍥剧墖浠嶅湪鐢熸垚涓紝鏆傛椂涓嶈兘鍒犻櫎鍘嗗彶璁板綍")
|
||||
}
|
||||
}
|
||||
if err := tx.Exec("UPDATE project_assets SET image_asset_id=NULL WHERE id=? AND image_asset_id=?", assetID, row.MediaID).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Exec("DELETE FROM generation_outputs WHERE id=?", outputID).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Exec("DELETE FROM generation_tasks WHERE id=?", row.TaskID).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Exec("DELETE FROM channel_asset_cache WHERE media_asset_id=?", row.MediaID).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Exec("DELETE FROM media_assets WHERE id=?", row.MediaID).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
objectKey = row.ObjectKey
|
||||
return nil
|
||||
})
|
||||
return objectKey, err
|
||||
}
|
||||
|
||||
func (s *Creative) UpdateEpisodeAnalysisSettings(userID, projectID, episodeID uuid.UUID, audioSource, sourceLanguage string) error {
|
||||
if audioSource != "video_audio" && audioSource != "subtitle_file" {
|
||||
return errors.New("音频来源无效")
|
||||
}
|
||||
sourceLanguage = strings.TrimSpace(sourceLanguage)
|
||||
if len(sourceLanguage) > 16 {
|
||||
return errors.New("源语言设置无效")
|
||||
}
|
||||
return s.DB.Transaction(func(tx *gorm.DB) error {
|
||||
var episode model.ProjectEpisode
|
||||
if err := tx.Table("project_episodes e").Select("e.*").
|
||||
Joins("JOIN creative_projects p ON p.id=e.project_id AND p.user_id=? AND p.deleted_at IS NULL", userID).
|
||||
Where("e.id=? AND e.project_id=? AND e.deleted_at IS NULL", episodeID, projectID).Take(&episode).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
// Selecting subtitle mode is persisted even before the file is uploaded.
|
||||
// The workbench disables analysis until a subtitle asset is present.
|
||||
updates := map[string]any{"audio_source": audioSource, "source_language": nil}
|
||||
if sourceLanguage != "" {
|
||||
updates["source_language"] = sourceLanguage
|
||||
}
|
||||
return tx.Model(&episode).Updates(updates).Error
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,108 @@
|
||||
// 本文件负责按创作项目类型筛选当前剧集可下载的视频,并生成 ZIP 文件清单。
|
||||
package service
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"path"
|
||||
"strings"
|
||||
|
||||
"juhe-factory/api/internal/videoarchive"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// episodeVideoArchiveRow 描述生成视频查询结果及其在分镜中的业务状态。
|
||||
type episodeVideoArchiveRow struct {
|
||||
ObjectKey string
|
||||
MimeType string
|
||||
SequenceNo int
|
||||
Active bool
|
||||
}
|
||||
|
||||
// EpisodeVideoArchive 查询当前剧集视频:短剧创作仅包含主视频和备选视频,视频转绘包含主视频和历史视频。
|
||||
func (s *Creative) EpisodeVideoArchive(userID, projectID, episodeID uuid.UUID) (videoarchive.Manifest, error) {
|
||||
// 校验项目归属并获取项目类型,当前剧集视频 ZIP 仅对短剧创作和视频转绘开放。
|
||||
var project struct {
|
||||
ProjectType string
|
||||
}
|
||||
if err := s.DB.Table("creative_projects").Select("project_type").
|
||||
Where("id=? AND user_id=? AND deleted_at IS NULL", projectID, userID).
|
||||
Take(&project).Error; err != nil {
|
||||
return videoarchive.Manifest{}, err
|
||||
}
|
||||
if project.ProjectType != "premium_drama" && project.ProjectType != "video_redraw" {
|
||||
return videoarchive.Manifest{}, gorm.ErrRecordNotFound
|
||||
}
|
||||
var episode struct {
|
||||
ProjectName string
|
||||
EpisodeName string
|
||||
}
|
||||
if err := s.DB.Table("project_episodes episode").
|
||||
Select("project.name AS project_name,episode.name AS episode_name").
|
||||
Joins("JOIN creative_projects project ON project.id=episode.project_id AND project.deleted_at IS NULL").
|
||||
Where("episode.id=? AND episode.project_id=? AND episode.deleted_at IS NULL", episodeID, projectID).
|
||||
Take(&episode).Error; err != nil {
|
||||
return videoarchive.Manifest{}, err
|
||||
}
|
||||
|
||||
rows := make([]episodeVideoArchiveRow, 0)
|
||||
query := s.DB.Table("generation_outputs output").
|
||||
Select(`media.object_key,media.mime_type,storyboard.sequence_no,
|
||||
storyboard.active_output_id=output.id AS active`).
|
||||
Joins("JOIN generation_tasks task ON task.id=output.task_id AND task.task_type='video_generation'").
|
||||
Joins("JOIN episode_storyboards storyboard ON storyboard.id=task.storyboard_id AND storyboard.deleted_at IS NULL").
|
||||
Joins("JOIN media_assets media ON media.id=output.media_asset_id AND media.deleted_at IS NULL").
|
||||
Where("storyboard.episode_id=?", episodeID)
|
||||
if project.ProjectType == "premium_drama" {
|
||||
query = query.Where("storyboard.active_output_id=output.id OR coalesce(output.metadata->>'candidate','false')='true'")
|
||||
}
|
||||
if err := query.Order("storyboard.sequence_no ASC,(storyboard.active_output_id=output.id) DESC,output.created_at DESC").Find(&rows).Error; err != nil {
|
||||
return videoarchive.Manifest{}, err
|
||||
}
|
||||
if len(rows) == 0 {
|
||||
return videoarchive.Manifest{}, errors.New("当前剧集暂无可下载视频")
|
||||
}
|
||||
|
||||
entries := make([]videoarchive.Entry, 0, len(rows))
|
||||
versions := make(map[int]int)
|
||||
for _, row := range rows {
|
||||
label := "主视频"
|
||||
if !row.Active {
|
||||
versions[row.SequenceNo]++
|
||||
if project.ProjectType == "premium_drama" {
|
||||
label = fmt.Sprintf("备选视频%d", versions[row.SequenceNo])
|
||||
} else {
|
||||
label = fmt.Sprintf("历史视频%d", versions[row.SequenceNo])
|
||||
}
|
||||
}
|
||||
name := fmt.Sprintf("分镜%03d-%s%s", row.SequenceNo, label, videoArchiveExtension(row.ObjectKey, row.MimeType))
|
||||
entries = append(entries, videoarchive.Entry{ObjectKey: row.ObjectKey, Name: videoarchive.SafeFilename(name, "视频.mp4")})
|
||||
}
|
||||
archiveName := videoarchive.SafeFilename(episode.ProjectName+episode.EpisodeName, "剧集视频") + ".zip"
|
||||
return videoarchive.Manifest{Name: archiveName, Entries: entries}, nil
|
||||
}
|
||||
|
||||
// videoArchiveExtension 根据已保存的媒体类型和对象键生成稳定的视频扩展名。
|
||||
func videoArchiveExtension(objectKey, mimeType string) string {
|
||||
switch strings.ToLower(strings.TrimSpace(strings.Split(mimeType, ";")[0])) {
|
||||
case "video/webm":
|
||||
return ".webm"
|
||||
case "video/quicktime":
|
||||
return ".mov"
|
||||
case "video/x-m4v":
|
||||
return ".m4v"
|
||||
case "video/x-matroska":
|
||||
return ".mkv"
|
||||
case "video/x-msvideo":
|
||||
return ".avi"
|
||||
case "video/mp4":
|
||||
return ".mp4"
|
||||
}
|
||||
extension := strings.ToLower(path.Ext(objectKey))
|
||||
if map[string]bool{".mp4": true, ".webm": true, ".mov": true, ".m4v": true, ".mkv": true, ".avi": true}[extension] {
|
||||
return extension
|
||||
}
|
||||
return ".mp4"
|
||||
}
|
||||
@@ -0,0 +1,544 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"archive/zip"
|
||||
"bytes"
|
||||
"crypto/sha256"
|
||||
"encoding/binary"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"encoding/xml"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"path/filepath"
|
||||
"regexp"
|
||||
"sort"
|
||||
"strings"
|
||||
"time"
|
||||
"unicode"
|
||||
"unicode/utf16"
|
||||
"unicode/utf8"
|
||||
|
||||
dramapkg "juhe-factory/api/internal/drama"
|
||||
"juhe-factory/api/internal/model"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/richardlehane/mscfb"
|
||||
"golang.org/x/text/encoding/simplifiedchinese"
|
||||
"golang.org/x/text/transform"
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/clause"
|
||||
)
|
||||
|
||||
const maxDramaImportBytes = 3 * 1024 * 1024
|
||||
|
||||
var chapterTitlePattern = regexp.MustCompile(`(?im)^\s*(?:#+\s*|[*+-]\s*)?((?:第[零〇一二三四五六七八九十百千万两\d]+[章节回卷篇部集幕].*)|(?:chapter\s+\d+.*))\s*$`)
|
||||
|
||||
type DramaChapter struct {
|
||||
Title string `json:"title"`
|
||||
Content string `json:"content,omitempty"`
|
||||
CharCount int `json:"char_count"`
|
||||
}
|
||||
|
||||
type DramaImportPreview struct {
|
||||
ImportToken uuid.UUID `json:"import_token"`
|
||||
HasChapters bool `json:"has_chapters"`
|
||||
TotalCharacters int `json:"total_characters"`
|
||||
Chapters []DramaChapter `json:"chapters"`
|
||||
SingleEpisodeConfirmationNeeded bool `json:"single_episode_confirmation_required"`
|
||||
}
|
||||
|
||||
func (s *Creative) PreviewDramaImport(userID, projectID uuid.UUID, filename string, data []byte, rawText string) (*DramaImportPreview, error) {
|
||||
if err := s.requirePremiumProject(userID, projectID); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(data) > maxDramaImportBytes || len([]byte(rawText)) > maxDramaImportBytes {
|
||||
return nil, errors.New("小说文件或文本不能超过3MB")
|
||||
}
|
||||
text, sourceType, err := parseDramaSource(filename, data, rawText)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
text = normalizeDramaText(text)
|
||||
if text == "" {
|
||||
return nil, errors.New("未读取到有效小说正文")
|
||||
}
|
||||
chapters := splitDramaChapters(text)
|
||||
metadata := make([]DramaChapter, 0, len(chapters))
|
||||
for _, chapter := range chapters {
|
||||
metadata = append(metadata, DramaChapter{Title: chapter.Title, CharCount: len([]rune(chapter.Content))})
|
||||
}
|
||||
previewJSON, _ := json.Marshal(map[string]any{"has_chapters": len(chapters) > 0, "chapters": metadata})
|
||||
hash := sha256.Sum256([]byte(text))
|
||||
session := model.DramaImportSession{
|
||||
ID: uuid.New(), UserID: userID, ProjectID: projectID, SourceType: sourceType,
|
||||
SourceFilename: cleanDramaFilename(filename), RawContent: text, ContentSHA256: hex.EncodeToString(hash[:]),
|
||||
PreviewData: previewJSON, ExpiresAt: time.Now().Add(30 * time.Minute),
|
||||
}
|
||||
if err := s.DB.Create(&session).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &DramaImportPreview{
|
||||
ImportToken: session.ID, HasChapters: len(chapters) > 0, TotalCharacters: len([]rune(text)),
|
||||
Chapters: metadata, SingleEpisodeConfirmationNeeded: len(chapters) == 0,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (s *Creative) ConfirmDramaImport(userID, projectID, token uuid.UUID, chaptersPerEpisode, startEpisodeNo int, treatAsSingle bool) ([]model.ProjectEpisode, error) {
|
||||
if chaptersPerEpisode < 1 {
|
||||
chaptersPerEpisode = 1
|
||||
}
|
||||
if startEpisodeNo < 1 {
|
||||
return nil, errors.New("起始集数无效")
|
||||
}
|
||||
created := make([]model.ProjectEpisode, 0)
|
||||
err := s.DB.Transaction(func(tx *gorm.DB) error {
|
||||
var project model.CreativeProject
|
||||
if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).Select("id").Where("id=? AND user_id=? AND deleted_at IS NULL", projectID, userID).Take(&project).Error; err != nil {
|
||||
return gorm.ErrRecordNotFound
|
||||
}
|
||||
var session model.DramaImportSession
|
||||
if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).Where("id=? AND user_id=? AND project_id=?", token, userID, projectID).Take(&session).Error; err != nil {
|
||||
return gorm.ErrRecordNotFound
|
||||
}
|
||||
if session.ConsumedAt != nil || time.Now().After(session.ExpiresAt) {
|
||||
return errors.New("导入预览已失效,请重新上传")
|
||||
}
|
||||
chapters := splitDramaChapters(session.RawContent)
|
||||
if len(chapters) == 0 {
|
||||
if !treatAsSingle {
|
||||
return errors.New("需要确认将全部内容作为一集导入")
|
||||
}
|
||||
chapters = []DramaChapter{{Title: fmt.Sprintf("第%d集", startEpisodeNo), Content: session.RawContent, CharCount: len([]rune(session.RawContent))}}
|
||||
chaptersPerEpisode = 1
|
||||
}
|
||||
for offset, index := 0, 0; index < len(chapters); offset, index = offset+1, index+chaptersPerEpisode {
|
||||
end := index + chaptersPerEpisode
|
||||
if end > len(chapters) {
|
||||
end = len(chapters)
|
||||
}
|
||||
group := chapters[index:end]
|
||||
episodeNo := startEpisodeNo + offset
|
||||
name := fmt.Sprintf("第%d集", episodeNo)
|
||||
if len(group) == 1 && strings.TrimSpace(group[0].Title) != "" {
|
||||
name = group[0].Title
|
||||
}
|
||||
name = truncateRunes(strings.TrimSpace(name), 160)
|
||||
exists, err := episodeNameExists(tx, projectID, nil, name)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if exists {
|
||||
return errors.New("同一项目下不可有同名剧集")
|
||||
}
|
||||
parts := make([]string, 0, len(group))
|
||||
for _, chapter := range group {
|
||||
parts = append(parts, strings.TrimSpace(chapter.Content))
|
||||
}
|
||||
content := strings.TrimSpace(strings.Join(parts, "\n\n"))
|
||||
if err := dramapkg.ValidateEpisodeContent(content); err != nil {
|
||||
return fmt.Errorf("第%d集:%w", episodeNo, err)
|
||||
}
|
||||
episode := model.ProjectEpisode{ID: uuid.New(), ProjectID: projectID, EpisodeNo: episodeNo, Name: name, AudioSource: "video_audio", Status: "draft"}
|
||||
if err := tx.Create(&episode).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
hash := sha256.Sum256([]byte(content))
|
||||
source := model.EpisodeSource{ID: uuid.New(), EpisodeID: episode.ID, Title: episode.Name, RawContent: content, CharCount: len([]rune(content)), SourceType: session.SourceType, SourceFilename: session.SourceFilename, ContentSHA256: hex.EncodeToString(hash[:])}
|
||||
if err := tx.Create(&source).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
created = append(created, episode)
|
||||
}
|
||||
now := time.Now()
|
||||
return tx.Model(&session).Updates(map[string]any{"consumed_at": now, "raw_content": ""}).Error
|
||||
})
|
||||
return created, err
|
||||
}
|
||||
|
||||
func (s *Creative) EpisodeSource(userID, projectID, episodeID uuid.UUID) (*model.EpisodeSource, error) {
|
||||
var source model.EpisodeSource
|
||||
result := s.DB.Table("episode_sources source").Joins("JOIN project_episodes episode ON episode.id=source.episode_id AND episode.deleted_at IS NULL").Joins("JOIN creative_projects project ON project.id=episode.project_id AND project.deleted_at IS NULL").Where("source.episode_id=? AND episode.project_id=? AND project.user_id=?", episodeID, projectID, userID).Limit(1).Find(&source)
|
||||
if result.Error != nil {
|
||||
return nil, result.Error
|
||||
}
|
||||
if result.RowsAffected == 0 {
|
||||
return nil, gorm.ErrRecordNotFound
|
||||
}
|
||||
return &source, nil
|
||||
}
|
||||
|
||||
func (s *Creative) SaveEpisodeSource(userID, projectID, episodeID uuid.UUID, content, title string) (*model.EpisodeSource, error) {
|
||||
content = normalizeDramaText(content)
|
||||
if err := dramapkg.ValidateEpisodeContent(content); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := s.requirePremiumEpisode(userID, projectID, episodeID); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
hash := sha256.Sum256([]byte(content))
|
||||
source := model.EpisodeSource{ID: uuid.New(), EpisodeID: episodeID, Title: truncateRunes(strings.TrimSpace(title), 200), RawContent: content, CharCount: len([]rune(content)), SourceType: "pasted_text", ContentSHA256: hex.EncodeToString(hash[:])}
|
||||
err := s.DB.Clauses(clause.OnConflict{Columns: []clause.Column{{Name: "episode_id"}}, DoUpdates: clause.AssignmentColumns([]string{"title", "raw_content", "char_count", "source_type", "source_filename", "content_sha256", "updated_at"})}).Create(&source).Error
|
||||
return &source, err
|
||||
}
|
||||
|
||||
func (s *Creative) requirePremiumProject(userID, projectID uuid.UUID) error {
|
||||
var count int64
|
||||
err := s.DB.Table("creative_projects").Where("id=? AND user_id=? AND project_type='premium_drama' AND deleted_at IS NULL", projectID, userID).Count(&count).Error
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if count == 0 {
|
||||
return gorm.ErrRecordNotFound
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *Creative) requirePremiumEpisode(userID, projectID, episodeID uuid.UUID) error {
|
||||
var count int64
|
||||
err := s.DB.Table("project_episodes episode").Joins("JOIN creative_projects project ON project.id=episode.project_id").Where("episode.id=? AND episode.project_id=? AND episode.deleted_at IS NULL AND project.user_id=? AND project.project_type='premium_drama' AND project.deleted_at IS NULL", episodeID, projectID, userID).Count(&count).Error
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if count == 0 {
|
||||
return gorm.ErrRecordNotFound
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func parseDramaSource(filename string, data []byte, rawText string) (string, string, error) {
|
||||
if strings.TrimSpace(rawText) != "" {
|
||||
if len(data) > 0 {
|
||||
return "", "", errors.New("文件和粘贴文本只能选择一种")
|
||||
}
|
||||
if looksLikeBinary([]byte(rawText)) {
|
||||
return "", "", errors.New("粘贴内容包含二进制控制字符")
|
||||
}
|
||||
if looksLikeHTMLOrScript([]byte(rawText)) {
|
||||
return "", "", errors.New("不支持HTML或脚本内容")
|
||||
}
|
||||
return rawText, "pasted_text", nil
|
||||
}
|
||||
if len(data) == 0 {
|
||||
return "", "", errors.New("请选择小说文件或粘贴文本")
|
||||
}
|
||||
ext := strings.ToLower(filepath.Ext(filename))
|
||||
switch ext {
|
||||
case ".txt":
|
||||
text, err := decodePlainText(data)
|
||||
return text, "txt", err
|
||||
case ".docx":
|
||||
text, err := extractDOCX(data)
|
||||
return text, "docx", err
|
||||
case ".doc":
|
||||
text, err := extractLegacyDOC(data)
|
||||
return text, "doc", err
|
||||
default:
|
||||
return "", "", errors.New("仅支持TXT、DOC和DOCX文件")
|
||||
}
|
||||
}
|
||||
|
||||
func decodePlainText(data []byte) (string, error) {
|
||||
if hasDisallowedBinarySignature(data) || bytes.IndexByte(data, 0) >= 0 || looksLikeBinary(data) {
|
||||
return "", errors.New("TXT文件包含二进制内容")
|
||||
}
|
||||
if looksLikeHTMLOrScript(data) {
|
||||
return "", errors.New("不支持HTML或脚本文件")
|
||||
}
|
||||
data = bytes.TrimPrefix(data, []byte{0xEF, 0xBB, 0xBF})
|
||||
if utf8.Valid(data) {
|
||||
return string(data), nil
|
||||
}
|
||||
decoded, err := io.ReadAll(transform.NewReader(bytes.NewReader(data), simplifiedchinese.GB18030.NewDecoder()))
|
||||
if err != nil || !utf8.Valid(decoded) {
|
||||
return "", errors.New("TXT文件编码不支持,请使用UTF-8或GB18030")
|
||||
}
|
||||
return string(decoded), nil
|
||||
}
|
||||
|
||||
func extractDOCX(data []byte) (string, error) {
|
||||
if len(data) < 4 || !bytes.Equal(data[:4], []byte{'P', 'K', 3, 4}) {
|
||||
return "", errors.New("DOCX文件结构无效")
|
||||
}
|
||||
reader, err := zip.NewReader(bytes.NewReader(data), int64(len(data)))
|
||||
if err != nil {
|
||||
return "", errors.New("DOCX文件结构无效")
|
||||
}
|
||||
if len(reader.File) > 1000 {
|
||||
return "", errors.New("DOCX文件包含过多内容")
|
||||
}
|
||||
var document, contentTypes []byte
|
||||
var expanded uint64
|
||||
for _, file := range reader.File {
|
||||
expanded += file.UncompressedSize64
|
||||
if expanded > 20*1024*1024 {
|
||||
return "", errors.New("DOCX展开后内容过大")
|
||||
}
|
||||
name := strings.ToLower(strings.ReplaceAll(file.Name, "\\", "/"))
|
||||
if strings.Contains(name, "vbaproject") || strings.Contains(name, "activex") || strings.Contains(name, "embeddings/") {
|
||||
return "", errors.New("DOCX包含宏、ActiveX或嵌入对象")
|
||||
}
|
||||
if name == "[content_types].xml" {
|
||||
stream, openErr := file.Open()
|
||||
if openErr != nil {
|
||||
return "", openErr
|
||||
}
|
||||
contentTypes, err = io.ReadAll(io.LimitReader(stream, 1024*1024))
|
||||
stream.Close()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
}
|
||||
if name == "word/document.xml" {
|
||||
stream, openErr := file.Open()
|
||||
if openErr != nil {
|
||||
return "", openErr
|
||||
}
|
||||
document, err = io.ReadAll(io.LimitReader(stream, 12*1024*1024))
|
||||
stream.Close()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
}
|
||||
}
|
||||
contentTypeText := strings.ToLower(string(contentTypes))
|
||||
if len(document) == 0 || !strings.Contains(contentTypeText, "wordprocessingml.document.main+xml") || strings.Contains(contentTypeText, "macroenabled") {
|
||||
return "", errors.New("DOCX缺少Word正文结构")
|
||||
}
|
||||
decoder := xml.NewDecoder(bytes.NewReader(document))
|
||||
var output strings.Builder
|
||||
for {
|
||||
token, tokenErr := decoder.Token()
|
||||
if tokenErr == io.EOF {
|
||||
break
|
||||
}
|
||||
if tokenErr != nil {
|
||||
return "", errors.New("DOCX正文XML无效")
|
||||
}
|
||||
switch value := token.(type) {
|
||||
case xml.CharData:
|
||||
output.Write([]byte(value))
|
||||
case xml.EndElement:
|
||||
if value.Name.Local == "p" {
|
||||
output.WriteString("\n")
|
||||
} else if value.Name.Local == "tab" {
|
||||
output.WriteString("\t")
|
||||
}
|
||||
}
|
||||
}
|
||||
return output.String(), nil
|
||||
}
|
||||
|
||||
func extractLegacyDOC(data []byte) (string, error) {
|
||||
magic := []byte{0xD0, 0xCF, 0x11, 0xE0, 0xA1, 0xB1, 0x1A, 0xE1}
|
||||
if len(data) < len(magic) || !bytes.Equal(data[:len(magic)], magic) {
|
||||
return "", errors.New("DOC文件结构无效")
|
||||
}
|
||||
reader, err := mscfb.New(bytes.NewReader(data))
|
||||
if err != nil {
|
||||
return "", errors.New("DOC文件结构无效")
|
||||
}
|
||||
streams := map[string][]byte{}
|
||||
for entry, nextErr := reader.Next(); nextErr == nil; entry, nextErr = reader.Next() {
|
||||
name := strings.ToLower(entry.Name)
|
||||
if strings.Contains(name, "vba") || strings.Contains(name, "macros") || strings.Contains(name, "objectpool") {
|
||||
return "", errors.New("DOC包含宏或嵌入对象")
|
||||
}
|
||||
if name == "worddocument" || name == "0table" || name == "1table" {
|
||||
content, readErr := io.ReadAll(io.LimitReader(entry, 12*1024*1024))
|
||||
if readErr != nil {
|
||||
return "", readErr
|
||||
}
|
||||
streams[name] = content
|
||||
}
|
||||
}
|
||||
word := streams["worddocument"]
|
||||
if len(word) < 0x1AA || binary.LittleEndian.Uint16(word[:2]) != 0xA5EC {
|
||||
return "", errors.New("DOC缺少有效WordDocument流")
|
||||
}
|
||||
flags := binary.LittleEndian.Uint16(word[0x0A:0x0C])
|
||||
tableName := "0table"
|
||||
if flags&0x0200 != 0 {
|
||||
tableName = "1table"
|
||||
}
|
||||
table := streams[tableName]
|
||||
fcClx := int(binary.LittleEndian.Uint32(word[0x1A2:0x1A6]))
|
||||
lcbClx := int(binary.LittleEndian.Uint32(word[0x1A6:0x1AA]))
|
||||
text, parseErr := extractDOCPieces(word, table, fcClx, lcbClx)
|
||||
if parseErr != nil || strings.TrimSpace(text) == "" {
|
||||
return "", errors.New("DOC正文无法安全解析")
|
||||
}
|
||||
return text, nil
|
||||
}
|
||||
|
||||
func extractDOCPieces(word, table []byte, offset, length int) (string, error) {
|
||||
if offset < 0 || length <= 0 || offset+length > len(table) {
|
||||
return "", errors.New("invalid CLX")
|
||||
}
|
||||
clx := table[offset : offset+length]
|
||||
pos := 0
|
||||
for pos < len(clx) && clx[pos] == 0x01 {
|
||||
if pos+3 > len(clx) {
|
||||
return "", io.ErrUnexpectedEOF
|
||||
}
|
||||
size := int(binary.LittleEndian.Uint16(clx[pos+1 : pos+3]))
|
||||
pos += 3 + size
|
||||
}
|
||||
if pos+5 > len(clx) || clx[pos] != 0x02 {
|
||||
return "", errors.New("missing Pcdt")
|
||||
}
|
||||
plcSize := int(binary.LittleEndian.Uint32(clx[pos+1 : pos+5]))
|
||||
pos += 5
|
||||
if plcSize < 4 || pos+plcSize > len(clx) {
|
||||
return "", errors.New("invalid PlcPcd")
|
||||
}
|
||||
plc := clx[pos : pos+plcSize]
|
||||
pieces := (plcSize - 4) / 12
|
||||
if pieces <= 0 || pieces > 100000 {
|
||||
return "", errors.New("invalid piece count")
|
||||
}
|
||||
cpBytes := (pieces + 1) * 4
|
||||
if cpBytes+pieces*8 > len(plc) {
|
||||
return "", io.ErrUnexpectedEOF
|
||||
}
|
||||
var output strings.Builder
|
||||
for i := 0; i < pieces; i++ {
|
||||
cpStart := int(binary.LittleEndian.Uint32(plc[i*4 : i*4+4]))
|
||||
cpEnd := int(binary.LittleEndian.Uint32(plc[(i+1)*4 : (i+1)*4+4]))
|
||||
if cpEnd <= cpStart {
|
||||
continue
|
||||
}
|
||||
pcd := cpBytes + i*8
|
||||
rawFC := binary.LittleEndian.Uint32(plc[pcd+2 : pcd+6])
|
||||
compressed := rawFC&0x40000000 != 0
|
||||
fc := int(rawFC & 0x3FFFFFFF)
|
||||
chars := cpEnd - cpStart
|
||||
if compressed {
|
||||
fc /= 2
|
||||
if fc < 0 || fc+chars > len(word) {
|
||||
continue
|
||||
}
|
||||
decoded, _ := io.ReadAll(transform.NewReader(bytes.NewReader(word[fc:fc+chars]), simplifiedchinese.GB18030.NewDecoder()))
|
||||
output.Write(decoded)
|
||||
} else {
|
||||
byteLen := chars * 2
|
||||
if fc < 0 || fc+byteLen > len(word) {
|
||||
continue
|
||||
}
|
||||
units := make([]uint16, chars)
|
||||
for j := 0; j < chars; j++ {
|
||||
units[j] = binary.LittleEndian.Uint16(word[fc+j*2 : fc+j*2+2])
|
||||
}
|
||||
output.WriteString(string(utf16.Decode(units)))
|
||||
}
|
||||
}
|
||||
return output.String(), nil
|
||||
}
|
||||
|
||||
func splitDramaChapters(content string) []DramaChapter {
|
||||
matches := chapterTitlePattern.FindAllStringSubmatchIndex(content, -1)
|
||||
if len(matches) == 0 {
|
||||
return nil
|
||||
}
|
||||
chapters := make([]DramaChapter, 0, len(matches))
|
||||
for i, match := range matches {
|
||||
start := match[0]
|
||||
if i == 0 {
|
||||
start = 0
|
||||
}
|
||||
end := len(content)
|
||||
if i+1 < len(matches) {
|
||||
end = matches[i+1][0]
|
||||
}
|
||||
body := strings.TrimSpace(content[start:end])
|
||||
if body == "" {
|
||||
continue
|
||||
}
|
||||
title := strings.TrimSpace(content[match[2]:match[3]])
|
||||
chapters = append(chapters, DramaChapter{Title: title, Content: body, CharCount: len([]rune(body))})
|
||||
}
|
||||
return chapters
|
||||
}
|
||||
|
||||
func normalizeDramaText(value string) string {
|
||||
value = strings.ReplaceAll(value, "\x00", "")
|
||||
value = strings.ReplaceAll(strings.ReplaceAll(value, "\r\n", "\n"), "\r", "\n")
|
||||
lines := strings.Split(value, "\n")
|
||||
cleaned := make([]string, 0, len(lines))
|
||||
blank := false
|
||||
for _, line := range lines {
|
||||
line = strings.TrimRightFunc(line, unicode.IsSpace)
|
||||
if strings.TrimSpace(line) == "" {
|
||||
if !blank {
|
||||
cleaned = append(cleaned, "")
|
||||
blank = true
|
||||
}
|
||||
continue
|
||||
}
|
||||
cleaned = append(cleaned, line)
|
||||
blank = false
|
||||
}
|
||||
return strings.TrimSpace(strings.Join(cleaned, "\n"))
|
||||
}
|
||||
|
||||
func looksLikeBinary(data []byte) bool {
|
||||
if len(data) == 0 {
|
||||
return false
|
||||
}
|
||||
controls := 0
|
||||
for _, b := range data {
|
||||
if b < 0x09 || (b > 0x0D && b < 0x20) {
|
||||
controls++
|
||||
}
|
||||
}
|
||||
return controls*100/len(data) > 2
|
||||
}
|
||||
func looksLikeHTMLOrScript(data []byte) bool {
|
||||
sample := strings.ToLower(string(data))
|
||||
markers := []string{"<!doctype html", "<html", "<head", "<body", "<script", "</script", "<iframe", "<object", "<embed", "javascript:", "vbscript:", "onerror=", "onload="}
|
||||
for _, marker := range markers {
|
||||
if strings.Contains(sample, marker) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
func hasDisallowedBinarySignature(data []byte) bool {
|
||||
signatures := [][]byte{
|
||||
{0x25, 0x50, 0x44, 0x46, 0x2D},
|
||||
{0x50, 0x4B, 0x03, 0x04},
|
||||
{0xD0, 0xCF, 0x11, 0xE0, 0xA1, 0xB1, 0x1A, 0xE1},
|
||||
{0x89, 0x50, 0x4E, 0x47},
|
||||
{0xFF, 0xD8, 0xFF},
|
||||
{0x47, 0x49, 0x46, 0x38},
|
||||
{0x4D, 0x5A},
|
||||
{0x7F, 0x45, 0x4C, 0x46},
|
||||
{0x52, 0x61, 0x72, 0x21},
|
||||
{0x37, 0x7A, 0xBC, 0xAF, 0x27, 0x1C},
|
||||
{0x1F, 0x8B},
|
||||
}
|
||||
data = bytes.TrimSpace(data)
|
||||
for _, signature := range signatures {
|
||||
if bytes.HasPrefix(data, signature) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
func cleanDramaFilename(value string) string {
|
||||
value = filepath.Base(strings.TrimSpace(value))
|
||||
return truncateRunes(value, 255)
|
||||
}
|
||||
func truncateRunes(value string, limit int) string {
|
||||
runes := []rune(value)
|
||||
if len(runes) > limit {
|
||||
return string(runes[:limit])
|
||||
}
|
||||
return value
|
||||
}
|
||||
|
||||
// Stable chapter ordering is useful in tests and when import previews are reconstructed.
|
||||
func sortEpisodesByNumber(items []model.ProjectEpisode) {
|
||||
sort.SliceStable(items, func(i, j int) bool { return items[i].EpisodeNo < items[j].EpisodeNo })
|
||||
}
|
||||
@@ -0,0 +1,95 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"archive/zip"
|
||||
"bytes"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestParseDramaSourceRejectsDisguisedAndExecutableContent(t *testing.T) {
|
||||
cases := []struct {
|
||||
name string
|
||||
filename string
|
||||
data []byte
|
||||
raw string
|
||||
}{
|
||||
{name: "pdf renamed txt", filename: "novel.txt", data: []byte("%PDF-1.7\ncontent")},
|
||||
{name: "zip renamed txt", filename: "novel.txt", data: []byte{'P', 'K', 3, 4, 1, 2, 3}},
|
||||
{name: "html after long prefix", filename: "novel.txt", data: []byte(strings.Repeat("正文", 3000) + "<script>alert(1)</script>")},
|
||||
{name: "script pasted", raw: "小说正文\njavascript:alert(1)"},
|
||||
{name: "binary pasted", raw: "正文\x00内容"},
|
||||
}
|
||||
for _, test := range cases {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
if _, _, err := parseDramaSource(test.filename, test.data, test.raw); err == nil {
|
||||
t.Fatal("expected unsafe import to be rejected")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtractDOCXReadsParagraphsAndValidatesWordPackage(t *testing.T) {
|
||||
valid := makeTestDOCX(t, `<?xml version="1.0"?><Types xmlns="http://schemas.openxmlformats.org/package/2006/content-types"><Override PartName="/word/document.xml" ContentType="application/vnd.openxmlformats-officedocument.wordprocessingml.document.main+xml"/></Types>`, nil)
|
||||
text, err := extractDOCX(valid)
|
||||
if err != nil {
|
||||
t.Fatalf("extractDOCX returned error: %v", err)
|
||||
}
|
||||
if !strings.Contains(text, "第一章 开始") || !strings.Contains(text, "这是正文") {
|
||||
t.Fatalf("unexpected extracted text: %q", text)
|
||||
}
|
||||
|
||||
fake := makeTestDOCX(t, `<?xml version="1.0"?><Types/>`, nil)
|
||||
if _, err := extractDOCX(fake); err == nil {
|
||||
t.Fatal("expected ZIP with fake DOCX structure to be rejected")
|
||||
}
|
||||
|
||||
macro := makeTestDOCX(t, `<?xml version="1.0"?><Types xmlns="http://schemas.openxmlformats.org/package/2006/content-types"><Override PartName="/word/document.xml" ContentType="application/vnd.openxmlformats-officedocument.wordprocessingml.document.main+xml"/></Types>`, map[string]string{"word/vbaProject.bin": "macro"})
|
||||
if _, err := extractDOCX(macro); err == nil {
|
||||
t.Fatal("expected macro-enabled package to be rejected")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSplitDramaChaptersPreservesOriginalText(t *testing.T) {
|
||||
source := "作品说明\n\n第一章 初见\n正文一\n\n第二章 重逢\n正文二"
|
||||
chapters := splitDramaChapters(source)
|
||||
if len(chapters) != 2 {
|
||||
t.Fatalf("expected 2 chapters, got %d", len(chapters))
|
||||
}
|
||||
joined := strings.TrimSpace(chapters[0].Content + "\n\n" + chapters[1].Content)
|
||||
if joined != source {
|
||||
t.Fatalf("chapter splitting lost or duplicated text:\nwant: %q\n got: %q", source, joined)
|
||||
}
|
||||
if chapters[0].Title != "第一章 初见" || chapters[1].Title != "第二章 重逢" {
|
||||
t.Fatalf("unexpected chapter titles: %#v", chapters)
|
||||
}
|
||||
if chapters := splitDramaChapters("没有章节标题的完整正文"); chapters != nil {
|
||||
t.Fatalf("text without chapter titles must require single-episode confirmation: %#v", chapters)
|
||||
}
|
||||
}
|
||||
|
||||
func makeTestDOCX(t *testing.T, contentTypes string, extras map[string]string) []byte {
|
||||
t.Helper()
|
||||
var buffer bytes.Buffer
|
||||
writer := zip.NewWriter(&buffer)
|
||||
files := map[string]string{
|
||||
"[Content_Types].xml": contentTypes,
|
||||
"word/document.xml": `<?xml version="1.0"?><w:document xmlns:w="http://schemas.openxmlformats.org/wordprocessingml/2006/main"><w:body><w:p><w:r><w:t>第一章 开始</w:t></w:r></w:p><w:p><w:r><w:t>这是正文</w:t></w:r></w:p></w:body></w:document>`,
|
||||
}
|
||||
for name, content := range extras {
|
||||
files[name] = content
|
||||
}
|
||||
for name, content := range files {
|
||||
entry, err := writer.Create(name)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err = entry.Write([]byte(content)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
if err := writer.Close(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return buffer.Bytes()
|
||||
}
|
||||
@@ -0,0 +1,232 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"time"
|
||||
|
||||
"juhe-factory/api/internal/billing"
|
||||
dramapkg "juhe-factory/api/internal/drama"
|
||||
"juhe-factory/api/internal/model"
|
||||
queuepkg "juhe-factory/api/internal/queue"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type dramaTextModel struct {
|
||||
ModelID uuid.UUID
|
||||
ChannelID uuid.UUID
|
||||
ModelName string
|
||||
ChannelName string
|
||||
Pricing billing.TextPricing `gorm:"-"`
|
||||
}
|
||||
|
||||
func (s *Creative) QueueDramaParse(userID, projectID, episodeID uuid.UUID) (*model.DramaParseTask, error) {
|
||||
if s.Queue == nil {
|
||||
return nil, errors.New("解析任务队列不可用")
|
||||
}
|
||||
if err := s.requirePremiumEpisode(userID, projectID, episodeID); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var source model.EpisodeSource
|
||||
if err := s.DB.Where("episode_id=?", episodeID).Take(&source).Error; err != nil {
|
||||
return nil, errors.New("请先导入或填写剧集原文")
|
||||
}
|
||||
if err := dramapkg.ValidateEpisodeContent(source.RawContent); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
configured, err := s.dramaTextModel(projectID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
prompt := dramaParsePrompt()
|
||||
snapshot, _ := json.Marshal(map[string]any{"model_id": configured.ModelID, "model_name": configured.ModelName, "channel_id": configured.ChannelID, "channel_name": configured.ChannelName})
|
||||
task := &model.DramaParseTask{
|
||||
ID: uuid.New(), RequestID: uuid.NewString(), UserID: userID, ProjectID: projectID,
|
||||
EpisodeID: episodeID, ChannelID: configured.ChannelID, ModelID: configured.ModelID,
|
||||
SourceSHA256: source.ContentSHA256, Status: "queued", ModelSnapshot: snapshot, PromptSnapshot: prompt,
|
||||
ContextSnapshot: json.RawMessage("{}"),
|
||||
BillingSnapshot: configured.Pricing.MarshalSnapshot(),
|
||||
EstimatedPoints: "0.00", PrepaidPoints: "0.00",
|
||||
}
|
||||
if err := s.DB.Transaction(func(tx *gorm.DB) error {
|
||||
var active int64
|
||||
if err := tx.Model(&model.GenerationTask{}).Where("episode_id=? AND status IN ?", episodeID, activeTaskStatuses).Count(&active).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if active > 0 {
|
||||
return errors.New("当前剧集存在进行中的生成任务,暂时不能重新解析")
|
||||
}
|
||||
if err := tx.Model(&model.DramaParseTask{}).Where("episode_id=? AND status IN ?", episodeID, []string{"queued", "running", "retry_wait", "cancel_requested"}).Count(&active).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if active > 0 {
|
||||
return errors.New("当前剧集已有解析任务")
|
||||
}
|
||||
if configured.Pricing.Mode == billing.TextBillingPerRequest {
|
||||
amount, err := billing.ChargeTextRequest(tx, userID, task.ID.String(), configured.Pricing, "剧本解析")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
task.EstimatedPoints = amount
|
||||
task.PrepaidPoints = amount
|
||||
}
|
||||
if err := tx.Omit("TokenCountSource").Create(task).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.Model(&model.ProjectEpisode{}).Where("id=?", episodeID).Updates(map[string]any{"status": "analyzing", "analysis_message": "剧本解析排队中"}).Error
|
||||
}); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := queuepkg.EnqueueID(s.Queue, queuepkg.TypeParseDramaEpisode, task.ID, 0); err != nil {
|
||||
s.DB.Model(task).Updates(map[string]any{"status": "failed", "error_code": "queue_unavailable", "error_message": err.Error(), "finished_at": time.Now()})
|
||||
s.DB.Model(&model.ProjectEpisode{}).Where("id=?", episodeID).Updates(map[string]any{"status": "failed", "analysis_message": "解析任务入队失败"})
|
||||
return nil, errors.New("解析任务入队失败")
|
||||
}
|
||||
return task, nil
|
||||
}
|
||||
|
||||
func (s *Creative) CancelDramaParse(userID, projectID, taskID uuid.UUID) error {
|
||||
result := s.DB.Model(&model.DramaParseTask{}).Where("id=? AND user_id=? AND project_id=? AND status IN ?", taskID, userID, projectID, []string{"queued", "running", "retry_wait"}).Updates(map[string]any{"status": "cancel_requested", "cancel_requested_at": time.Now(), "error_code": "user_cancelled", "error_message": "用户主动取消,费用不退"})
|
||||
if result.Error != nil {
|
||||
return result.Error
|
||||
}
|
||||
if result.RowsAffected == 0 {
|
||||
return gorm.ErrRecordNotFound
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *Creative) ListDramaParseTasks(userID, projectID uuid.UUID, episodeID *uuid.UUID) ([]model.DramaParseTask, error) {
|
||||
items := make([]model.DramaParseTask, 0)
|
||||
query := s.DB.Where("user_id=? AND project_id=?", userID, projectID)
|
||||
if episodeID != nil {
|
||||
query = query.Where("episode_id=?", *episodeID)
|
||||
}
|
||||
err := query.Order("created_at DESC").Limit(100).Find(&items).Error
|
||||
return items, err
|
||||
}
|
||||
|
||||
func (s *Creative) dramaTextModel(projectID uuid.UUID) (dramaTextModel, error) {
|
||||
var item dramaTextModel
|
||||
err := s.DB.Table("project_model_configs config").Select("model.id AS model_id,model.channel_id,model.name AS model_name,channel.name AS channel_name").Joins("JOIN models model ON model.id=config.model_id AND model.model_type='text' AND model.enabled=true AND model.deleted_at IS NULL").Joins("JOIN channels channel ON channel.id=model.channel_id AND channel.enabled=true AND channel.deleted_at IS NULL").Where("config.project_id=? AND config.model_type='text'", projectID).Take(&item).Error
|
||||
if err != nil {
|
||||
return item, errors.New("请先配置短剧创作文本模型")
|
||||
}
|
||||
item.Pricing, err = billing.LoadTextPricing(s.DB, item.ModelID)
|
||||
if err != nil {
|
||||
return item, errors.New("短剧创作文本模型计费配置无效")
|
||||
}
|
||||
return item, nil
|
||||
}
|
||||
|
||||
func dramaParsePrompt() string {
|
||||
return "【剧本解析规则】\n" + dramaStoryboardRules + "\n\n【角色、场景、道具解析规则】\n" + dramaEntityRules + "\n\n" + dramaOutputContract
|
||||
}
|
||||
|
||||
const dramaStoryboardRules = `你是一名资深影视剧本分析师和分镜设计师。请结合已有角色、场景、道具资料和当前剧集原文完成短剧分镜分析。
|
||||
|
||||
【层级定义】
|
||||
1. 一个 storyboard 是工作台中的一个“分镜”,对应一段完整的 5 至 15 秒视频。
|
||||
2. 一个 storyboard 内部可以包含多个摄影镜头。摄影镜头是该段视频中的机位、景别或画面切换,不得把每个摄影镜头分别输出为独立 storyboard。
|
||||
3. duration_seconds 表示整个 storyboard 的总时长,不是内部某个镜头的时长,必须是 5 至 15 之间的整数。
|
||||
4. script_content 必须按时间顺序列出该分镜内的全部摄影镜头。各镜头时间必须连续,不得重叠或断档,时长之和必须等于 duration_seconds。
|
||||
|
||||
【基本要求】
|
||||
1. 只处理本次提供的原文或分段,不得自行补写、删改剧情,绝对禁止修改角色台词。
|
||||
2. 按叙事顺序完整覆盖原文,不得遗漏关键动作、对白和剧情转折。
|
||||
3. 根据完整动作阶段和叙事意义划分 storyboard,不得为了凑时长合并无关剧情,也不得把一个连续动作机械拆成多个 storyboard。
|
||||
4. 分镜中的角色、场景和道具必须使用实体列表中的正式名称,不得使用别名、泛称或未定义名称。
|
||||
5. script_content 是工作台直接展示的完整分镜正文,必须使用下方“分镜正文固定格式”,不得只写剧情概述。
|
||||
6. prompt_content 是与整个 storyboard 对应、可直接用于视频生成的完整视觉提示词,也必须使用下方“分镜正文固定格式”,不得省略第一帧、站位或镜头细节。
|
||||
7. dialogue 提取本分镜原文中实际存在的对白;没有对白时返回空数组。
|
||||
8. asset_names 只列出本分镜实际出现的实体正式名称,并按 characters、scenes、props 分类。
|
||||
9. source_excerpt 必须引用本分镜所依据的当前原文片段。
|
||||
|
||||
【分镜正文固定格式】
|
||||
script_content 和 prompt_content 必须分别输出完整内容,并严格使用以下结构:
|
||||
场景=场景资产正式名称
|
||||
【第一帧】
|
||||
- 画面:明确景别、构图、角色外观、服装、表情、环境和光线色调。
|
||||
- 站位:明确每个角色及关键道具在画面中的前后左右关系、朝向、视线和接触关系。
|
||||
【画面内容】
|
||||
【镜头 1】(0-X秒)
|
||||
- 镜头:明确景别、机位角度、构图和运镜方式。
|
||||
- 画面:写清角色站位、连续动作、表情与视线变化、环境变化、光影色调、关键道具交互,以及镜头结束时的画面。
|
||||
- 声音/台词:写入该时间段实际发生的环境声、动作声和原文对白;没有则写“无”。
|
||||
【镜头 2】(X-duration_seconds秒)
|
||||
- 按相同项目继续描述,直至覆盖整个分镜。
|
||||
N 使用该分镜在返回数组中的顺序,从 1 开始;所有镜头时间连续、不重叠、不留空,最后一个镜头结束时间必须等于 duration_seconds。台词必须写入实际发生的镜头段落,禁止在末尾集中罗列。
|
||||
禁止只写“保持统一风格”“自然运镜”“角色互动”“气氛紧张”等空泛描述;必须提供具体、可拍摄、可执行的视觉和声音细节。
|
||||
|
||||
【内部镜头规则】
|
||||
1. 每个内部镜头按照“分镜正文固定格式”独立分段,并填写镜头、画面、声音/台词三项。
|
||||
2. 每个镜头必须描述开始时的主体状态、镜头内连续动作、人物视线或表情变化,以及结束时画面停留的位置。
|
||||
3. 关键动作、关键道具、信息揭示和人物反应应通过不同镜头清楚呈现;明显改变景别、机位、视角或主体时,应划分新的内部镜头。
|
||||
4. 5 至 7 秒通常包含 2 至 3 个镜头,8 至 12 秒通常包含 3 至 5 个镜头,13 至 15 秒通常包含 4 至 6 个镜头。以清楚呈现原文为准,不得无意义切镜。
|
||||
5. 原文明确适合长镜头时可以只有一个镜头,但必须写清持续动作、运镜过程和画面变化,不能只写静态概述。
|
||||
6. 只能描述观众能够看到或听到的内容。不得用“人物震惊”“气氛紧张”“陷入回忆”等抽象概括代替可见的表情、动作、声音和画面变化。
|
||||
7. 保持人物位置、服装、视线方向、道具状态、时间和环境连续;任何变化都必须有原文依据或在镜头内交代。
|
||||
8. 内部镜头之间按需标明硬切、视线匹配、动作匹配、声音先入等衔接关系。
|
||||
|
||||
【输出前自检】
|
||||
1. 每个 storyboard 是否为 5 至 15 秒。
|
||||
2. 内部镜头是否按时间连续排列,最后一个镜头结束时间是否等于 duration_seconds。
|
||||
3. 是否把内部摄影镜头错误拆成多个 storyboard,或把无关剧情错误合并为一个 storyboard。
|
||||
4. 每个内部镜头是否具有时间段、景别、机位、运镜和可见动作。
|
||||
5. 关键动作、线索特写、人物反应和原文对白是否完整覆盖。
|
||||
6. 是否存在抽象心理、笼统气氛或无法拍摄的概括;如有,改成可见或可听内容。
|
||||
7. 角色、场景、道具名称和连续性是否正确。
|
||||
完成自检后,按平台固定协议返回结果,不输出自检过程。`
|
||||
|
||||
const dramaEntityRules = `必须按以下规则识别和更新角色、场景、道具。
|
||||
|
||||
【通用规则】
|
||||
1. 先与用户提供的已有实体比对。相同实体必须复用已有正式名称,不得重复创建。
|
||||
2. 原文中的昵称、职位称呼、简称、化名和不同写法应合并到同一实体的 aliases;不得把别名当成新实体。
|
||||
3. 仅提取原文明示或可直接、客观推断的实体。旁白、叙述者、无明确身份的群演不作为角色资产。
|
||||
4. 已有实体原则上保留;当前原文明确提供新增或变化信息时才更新。不得用推测覆盖用户已经确认的资料。
|
||||
5. 同一角色只有在原文明示明显不同的时代、年龄阶段或造型,且需要独立视觉资产时,才拆成带限定词的独立正式名称。
|
||||
6. canonical_name 使用简洁、稳定的中文正式名称;aliases 去重且不得包含 canonical_name。
|
||||
7. description 写适合资产列表展示的一至两句客观视觉摘要。详细信息必须写入 attributes,避免只放在 description 中。
|
||||
|
||||
【角色 characters】
|
||||
attributes 尽量使用以下键:age、gender、identity、appearance、hairstyle、costume。外貌需描述可见且相对稳定的体型、脸型、肤色和五官特征;发型需描述长度、颜色和样式;服装需描述上装、下装、鞋履的颜色、款式和材质。禁止把情绪、瞬时表情、动作或性格当作固定外观。
|
||||
|
||||
【场景 scenes】
|
||||
attributes 尽量使用以下键:space_type、location、structure、visual_features、key_props。描述室内或室外、地理位置、空间结构、建筑或装饰风格、主要材质、色调、光线及固定陈设。不要把同一地点的不同称呼拆成多个场景;只有空间本身明显不同才拆分。
|
||||
|
||||
【道具 props】
|
||||
attributes 尽量使用以下键:type、appearance、function。描述道具类别、材质、颜色、尺寸、外形,以及原文明确体现的功能或剧情作用。普通背景陈设若不影响剧情且不需要单独生成,不作为独立道具。
|
||||
|
||||
【一致性检查】
|
||||
最终输出前检查:每个分镜 asset_names 中的名称都能在 entities 对应数组或用户提供的已有实体中找到;分镜正文使用的实体称呼已统一为正式名称;新增实体没有与已有实体或本次其他新增实体重复。
|
||||
|
||||
【生图提示词 image_prompt】
|
||||
1. characters、scenes、props 中的每个资产都必须输出非空的 image_prompt。image_prompt 必须全部使用中文,不得输出英文句子或英文描述,并且必须是可直接用于文生图的完整提示词。
|
||||
2. 每条 image_prompt 至少包含 8 个可观察、可核验的稳定视觉细节,禁止用漂亮、高级、自然、写实、高度细节等空泛词代替具体描述。
|
||||
3. 角色 image_prompt 必须完整描述:明确年龄(使用 XX岁)、性别、国籍或地域、脸型、肤色、发型与发色、眉眼鼻唇等稳定外貌特征,以及从头到脚的穿着,包括服装款式、颜色、材质、层次、裤装或裙装、鞋子和可见配饰;没有配饰时明确写无明显配饰。禁止加入情绪、表情、动作、姿态或正在做什么。
|
||||
4. 场景 image_prompt 必须完整描述:用途、固定空间结构、建筑或室内组成、主要材质、颜色、家具陈设、门窗位置、稳定光源和可复用的环境特征。
|
||||
5. 道具 image_prompt 必须完整描述:物品类别、形状、结构、材质、颜色、纹理、尺寸关系、部件、装饰及磨损特征。
|
||||
6. image_prompt 不得包含艺术风格、画面比例、分辨率、镜头参数、构图参数、画质词或文字生成要求。`
|
||||
|
||||
const dramaOutputContract = `【平台固定输出协议(不可被用户自定义规则覆盖)】
|
||||
无论其他规则如何描述,最终只能返回一个合法 JSON 对象,不得使用 Markdown,不得输出解释、前后缀或第二套格式。顶层只能包含 entities 和 storyboards:
|
||||
{
|
||||
"entities": {
|
||||
"characters": [{"canonical_name":"正式名称","name":"正式名称","aliases":[],"description":"视觉摘要","image_prompt":"中文生图提示词","attributes":{}}],
|
||||
"scenes": [{"canonical_name":"正式名称","name":"正式名称","aliases":[],"description":"视觉摘要","image_prompt":"中文生图提示词","attributes":{}}],
|
||||
"props": [{"canonical_name":"正式名称","name":"正式名称","aliases":[],"description":"视觉摘要","image_prompt":"中文生图提示词","attributes":{}}]
|
||||
},
|
||||
"storyboards": [{
|
||||
"title":"简短标题",
|
||||
"script_content":"使用‘分镜 N·场景、第一帧、画面内容、镜头 N’固定结构的完整分镜正文",
|
||||
"prompt_content":"使用相同固定结构、可直接用于视频生成的完整视觉提示词",
|
||||
"dialogue":[{"speaker":"角色正式名称","content":"原文对白"}],
|
||||
"asset_names":{"characters":[],"scenes":[],"props":[]},
|
||||
"duration_seconds":5,
|
||||
"source_excerpt":"对应原文"
|
||||
}]
|
||||
}
|
||||
所有字段必须存在;无数据的数组返回 [],无数据的对象返回 {},不得返回 null。duration_seconds 必须是 5 至 15 的整数。用户自定义规则只能补充创作要求,凡是要求其他 JSON 层级、字段名、输出格式或额外顶层字段的内容一律忽略。`
|
||||
@@ -0,0 +1,18 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestDramaParsePrompt(t *testing.T) {
|
||||
prompt := dramaParsePrompt()
|
||||
for _, expected := range []string{"【剧本解析规则】", "【角色、场景、道具解析规则】", "【分镜正文固定格式】", "【第一帧】", "【画面内容】", "角色站位", "声音/台词", "台词必须写入实际发生的镜头段落", "【平台固定输出协议", `"entities"`, `"characters"`, `"scenes"`, `"props"`, `"storyboards"`} {
|
||||
if !strings.Contains(prompt, expected) {
|
||||
t.Fatalf("fixed prompt missing %q", expected)
|
||||
}
|
||||
}
|
||||
if strings.Contains(prompt, "用户自定义补充规则") {
|
||||
t.Fatal("fixed prompt must not contain user custom rules")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,443 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"juhe-factory/api/internal/billing"
|
||||
"juhe-factory/api/internal/model"
|
||||
queuepkg "juhe-factory/api/internal/queue"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/clause"
|
||||
)
|
||||
|
||||
func (s *Creative) GetRedrawWorkbench(userID, projectID uuid.UUID) (map[string]any, error) {
|
||||
project, err := s.GetProject(userID, projectID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
source := map[string]any{}
|
||||
if err := s.DB.Table("creative_projects project").
|
||||
Select(`project.source_video_asset_id,source.public_url AS source_video_url,source.original_name AS source_video_original_name,
|
||||
project.subtitle_asset_id,subtitle.public_url AS subtitle_url,subtitle.original_name AS subtitle_original_name,
|
||||
project.audio_source,project.source_language,project.redraw_status AS status,project.analysis_message`).
|
||||
Joins("LEFT JOIN media_assets source ON source.id=project.source_video_asset_id AND source.deleted_at IS NULL").
|
||||
Joins("LEFT JOIN media_assets subtitle ON subtitle.id=project.subtitle_asset_id AND subtitle.deleted_at IS NULL").
|
||||
Where("project.id=? AND project.user_id=? AND project.project_type='video_redraw' AND project.deleted_at IS NULL", projectID, userID).
|
||||
Take(&source).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
assets := make([]map[string]any, 0)
|
||||
if err := s.DB.Table("project_assets asset").
|
||||
Select("asset.id,asset.project_id,asset.asset_type,asset.name,asset.user_edited").
|
||||
Where("asset.project_id=? AND asset.deleted_at IS NULL", projectID).
|
||||
Order("asset.asset_type,asset.created_at").Find(&assets).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
storyboards := make([]map[string]any, 0)
|
||||
if err := s.DB.Table("episode_storyboards storyboard").
|
||||
Select(`storyboard.*,thumbnail.public_url AS thumbnail_url,0 AS history_count`).
|
||||
Joins("LEFT JOIN media_assets thumbnail ON thumbnail.id=storyboard.thumbnail_asset_id").
|
||||
Where("storyboard.project_id=? AND storyboard.deleted_at IS NULL", projectID).
|
||||
Order("storyboard.sequence_no").Find(&storyboards).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var scriptRow struct {
|
||||
RedrawScript string
|
||||
}
|
||||
if err := s.DB.Table("creative_projects").Select("coalesce(redraw_script,'') AS redraw_script").
|
||||
Where("id=? AND user_id=? AND project_type='video_redraw' AND deleted_at IS NULL", projectID, userID).
|
||||
Take(&scriptRow).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return map[string]any{"project": project, "source": source, "assets": assets, "storyboards": storyboards, "script_content": scriptRow.RedrawScript}, nil
|
||||
}
|
||||
|
||||
func (s *Creative) GetRedrawEpisodeWorkbench(userID, projectID, episodeID uuid.UUID) (map[string]any, error) {
|
||||
data, err := s.GetWorkbench(userID, projectID, episodeID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
episode, ok := data["episode"].(map[string]any)
|
||||
if !ok {
|
||||
return nil, errors.New("剧集数据无效")
|
||||
}
|
||||
var script string
|
||||
if err := s.DB.Table("project_episodes episode").
|
||||
Select("coalesce(episode.redraw_script,'')").
|
||||
Joins("JOIN creative_projects project ON project.id=episode.project_id AND project.user_id=? AND project.project_type='video_redraw' AND project.deleted_at IS NULL", userID).
|
||||
Where("episode.id=? AND episode.project_id=? AND episode.deleted_at IS NULL", episodeID, projectID).
|
||||
Scan(&script).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return map[string]any{
|
||||
"project": data["project"], "episode": episode, "source": episode,
|
||||
"assets": data["assets"], "storyboards": data["storyboards"], "script_content": script,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// SaveRedrawEpisodeScript 校验项目归属和类型后保存用户编辑的剧集反推剧本。
|
||||
func (s *Creative) SaveRedrawEpisodeScript(userID, projectID, episodeID uuid.UUID, content string) error {
|
||||
return s.DB.Transaction(func(tx *gorm.DB) error {
|
||||
var episode model.ProjectEpisode
|
||||
if err := tx.Table("project_episodes episode").Select("episode.id").
|
||||
Joins("JOIN creative_projects project ON project.id=episode.project_id AND project.user_id=? AND project.project_type='video_redraw' AND project.deleted_at IS NULL", userID).
|
||||
Where("episode.id=? AND episode.project_id=? AND episode.deleted_at IS NULL", episodeID, projectID).
|
||||
Take(&episode).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.Model(&model.ProjectEpisode{}).Where("id=?", episode.ID).Update("redraw_script", content).Error
|
||||
})
|
||||
}
|
||||
|
||||
func (s *Creative) QueueRedrawEpisodeScript(userID, projectID, episodeID uuid.UUID) (*model.GenerationTask, error) {
|
||||
if s.Queue == nil {
|
||||
return nil, errors.New("反推任务队列不可用")
|
||||
}
|
||||
config, err := s.analysisModel(projectID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
task := &model.GenerationTask{
|
||||
RequestID: "script_reverse_" + uuid.NewString(), UserID: userID, ChannelID: &config.ChannelID,
|
||||
ModelID: &config.ModelID, ProjectID: &projectID, EpisodeID: &episodeID, TaskType: "prompt_reverse", Status: "submitted",
|
||||
InputData: json.RawMessage(`{"mode":"script"}`), EstimatedPoints: "0.00", PrepaidPoints: "0.00",
|
||||
}
|
||||
if err := s.DB.Transaction(func(tx *gorm.DB) error {
|
||||
scriptPrompt, err := selectedPromptContent(tx, userID, "剧本反推")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
task.InputData = mustJSON(map[string]string{"mode": "script", "prompt": scriptPrompt})
|
||||
var episode model.ProjectEpisode
|
||||
if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).Table("project_episodes episode").Select("episode.*").
|
||||
Joins("JOIN creative_projects project ON project.id=episode.project_id AND project.user_id=? AND project.project_type='video_redraw' AND project.deleted_at IS NULL", userID).
|
||||
Where("episode.id=? AND episode.project_id=? AND episode.deleted_at IS NULL", episodeID, projectID).Take(&episode).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if episode.Status != "review" {
|
||||
return errors.New("请等待全部分镜提示词反推完成")
|
||||
}
|
||||
var total, incomplete int64
|
||||
if err := tx.Model(&model.EpisodeStoryboard{}).Where("episode_id=? AND deleted_at IS NULL", episodeID).Count(&total).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if total == 0 {
|
||||
return errors.New("暂无可用于反推剧本的分镜提示词")
|
||||
}
|
||||
if err := tx.Model(&model.EpisodeStoryboard{}).
|
||||
Where("episode_id=? AND deleted_at IS NULL AND trim(coalesce(prompt_content,''))=''", episodeID).
|
||||
Count(&incomplete).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if incomplete > 0 {
|
||||
return errors.New("请等待所有分镜提示词反推完成")
|
||||
}
|
||||
var prompts []struct {
|
||||
SequenceNo int `json:"sequence_no"`
|
||||
StartMS int64 `json:"start_ms"`
|
||||
EndMS int64 `json:"end_ms"`
|
||||
PromptContent string `json:"prompt_content"`
|
||||
}
|
||||
if err := tx.Model(&model.EpisodeStoryboard{}).Select("sequence_no,start_ms,end_ms,prompt_content").
|
||||
Where("episode_id=? AND deleted_at IS NULL", episodeID).Order("sequence_no").Find(&prompts).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
task.PromptSnapshot = string(mustJSON(prompts))
|
||||
var active int64
|
||||
if err := tx.Model(&model.GenerationTask{}).Where("episode_id=? AND task_type='prompt_reverse' AND status IN ?", episodeID, activeTaskStatuses).Count(&active).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if active > 0 {
|
||||
return errors.New("当前剧集已有反推任务正在执行")
|
||||
}
|
||||
if err := billing.CreateTextGenerationTask(tx, task, config.Pricing, "剧本反推"); err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.Model(&episode).Updates(map[string]any{"status": "generating", "analysis_message": "正在反推完整剧本"}).Error
|
||||
}); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := queuepkg.EnqueueID(s.Queue, queuepkg.TypeAnalyzeEpisode, task.ID, 0); err != nil {
|
||||
if settleErr := s.failQueuedTask(task.ID, "queue_unavailable", "剧本反推任务入队失败"); settleErr != nil {
|
||||
return nil, fmt.Errorf("剧本反推任务入队失败且预扣返还失败: %w", settleErr)
|
||||
}
|
||||
_ = s.DB.Model(&model.ProjectEpisode{}).Where("id=?", episodeID).Updates(map[string]any{"status": "review", "analysis_message": "剧本反推任务入队失败"}).Error
|
||||
return nil, errors.New("反推任务队列暂时不可用,请稍后重试")
|
||||
}
|
||||
return task, nil
|
||||
}
|
||||
|
||||
func mustJSON(value any) []byte {
|
||||
result, _ := json.Marshal(value)
|
||||
return result
|
||||
}
|
||||
|
||||
// selectedPromptContent 按用户偏好选取指定类型的提示词内容,未命中用户偏好时兜底取最新系统提示词。
|
||||
// 用户未设置偏好或系统未配置均属于正常业务情况,使用 Limit(1).Find 避免触发 GORM 的 record not found 日志。
|
||||
func selectedPromptContent(tx *gorm.DB, userID uuid.UUID, promptType string) (string, error) {
|
||||
var selected struct{ Content string }
|
||||
if err := tx.Table("prompts p").Select("p.content").
|
||||
Joins("JOIN user_prompt_preferences pref ON pref.prompt_id=p.id AND pref.user_id=? AND pref.prompt_type=?", userID, promptType).
|
||||
Where("p.type=? AND p.deleted_at IS NULL AND (p.scope='system' OR (p.scope='user' AND p.owner_user_id=?))", promptType, userID).
|
||||
Limit(1).Find(&selected).Error; err != nil {
|
||||
return "", err
|
||||
}
|
||||
if strings.TrimSpace(selected.Content) == "" {
|
||||
if err := tx.Table("prompts").Select("content").
|
||||
Where("type=? AND scope='system' AND deleted_at IS NULL", promptType).
|
||||
Order("updated_at DESC").Limit(1).Find(&selected).Error; err != nil {
|
||||
return "", err
|
||||
}
|
||||
if strings.TrimSpace(selected.Content) == "" {
|
||||
return "", fmt.Errorf("请先在提示词管理中配置%s提示词", promptType)
|
||||
}
|
||||
}
|
||||
return strings.TrimSpace(selected.Content), nil
|
||||
}
|
||||
|
||||
func (s *Creative) UpdateRedrawAnalysisSettings(userID, projectID uuid.UUID, audioSource, sourceLanguage string) error {
|
||||
if audioSource != "video_audio" && audioSource != "subtitle_file" {
|
||||
return errors.New("台词来源无效")
|
||||
}
|
||||
sourceLanguage = strings.TrimSpace(sourceLanguage)
|
||||
var language any
|
||||
if sourceLanguage != "" {
|
||||
language = sourceLanguage
|
||||
}
|
||||
result := s.DB.Table("creative_projects").
|
||||
Where("id=? AND user_id=? AND project_type='video_redraw' AND deleted_at IS NULL", projectID, userID).
|
||||
Updates(map[string]any{"audio_source": audioSource, "source_language": language})
|
||||
if result.Error != nil {
|
||||
return result.Error
|
||||
}
|
||||
if result.RowsAffected == 0 {
|
||||
return gorm.ErrRecordNotFound
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *Creative) QueueRedrawAnalysis(userID, projectID uuid.UUID) (*model.GenerationTask, error) {
|
||||
return s.queueRedrawAnalysis(userID, projectID, nil)
|
||||
}
|
||||
|
||||
func (s *Creative) QueueRedrawStoryboardAnalysis(userID, projectID, storyboardID uuid.UUID) (*model.GenerationTask, error) {
|
||||
return s.queueRedrawAnalysis(userID, projectID, &storyboardID)
|
||||
}
|
||||
|
||||
func (s *Creative) QueueRedrawScript(userID, projectID uuid.UUID) (*model.GenerationTask, error) {
|
||||
if s.Queue == nil {
|
||||
return nil, errors.New("反推任务队列不可用")
|
||||
}
|
||||
config, err := s.analysisModel(projectID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
task := &model.GenerationTask{
|
||||
RequestID: "script_reverse_" + uuid.NewString(), UserID: userID, ChannelID: &config.ChannelID,
|
||||
ModelID: &config.ModelID, ProjectID: &projectID, TaskType: "prompt_reverse", Status: "submitted",
|
||||
InputData: json.RawMessage(`{"mode":"script"}`), EstimatedPoints: "0.00", PrepaidPoints: "0.00",
|
||||
}
|
||||
if err := s.DB.Transaction(func(tx *gorm.DB) error {
|
||||
scriptPrompt, err := selectedPromptContent(tx, userID, "剧本反推")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
task.InputData = mustJSON(map[string]string{"mode": "script", "prompt": scriptPrompt})
|
||||
var project model.CreativeProject
|
||||
if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).
|
||||
Where("id=? AND user_id=? AND project_type='video_redraw' AND deleted_at IS NULL", projectID, userID).
|
||||
Take(&project).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if project.RedrawStatus != "review" {
|
||||
return errors.New("请等待全部分镜提示词反推完成")
|
||||
}
|
||||
var total, incomplete int64
|
||||
if err := tx.Model(&model.EpisodeStoryboard{}).Where("project_id=? AND deleted_at IS NULL", projectID).Count(&total).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if total == 0 {
|
||||
return errors.New("暂无可用于反推剧本的分镜提示词")
|
||||
}
|
||||
if err := tx.Model(&model.EpisodeStoryboard{}).
|
||||
Where("project_id=? AND deleted_at IS NULL AND trim(coalesce(prompt_content,''))=''", projectID).
|
||||
Count(&incomplete).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if incomplete > 0 {
|
||||
return errors.New("请等待所有分镜提示词反推完成")
|
||||
}
|
||||
var storyboardPrompts []struct {
|
||||
SequenceNo int `json:"sequence_no"`
|
||||
StartMS int64 `json:"start_ms"`
|
||||
EndMS int64 `json:"end_ms"`
|
||||
PromptContent string `json:"prompt_content"`
|
||||
}
|
||||
if err := tx.Model(&model.EpisodeStoryboard{}).Select("sequence_no,start_ms,end_ms,prompt_content").
|
||||
Where("project_id=? AND deleted_at IS NULL", projectID).Order("sequence_no").Find(&storyboardPrompts).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
promptSnapshot, err := json.Marshal(storyboardPrompts)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
task.PromptSnapshot = string(promptSnapshot)
|
||||
var active int64
|
||||
if err := tx.Model(&model.GenerationTask{}).Where("project_id=? AND task_type='prompt_reverse' AND status IN ?", projectID, activeTaskStatuses).Count(&active).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if active > 0 {
|
||||
return errors.New("当前剧本已有反推任务正在执行")
|
||||
}
|
||||
if err := billing.CreateTextGenerationTask(tx, task, config.Pricing, "剧本反推"); err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.Model(&project).Updates(map[string]any{"redraw_status": "generating", "analysis_message": "正在反推完整剧本"}).Error
|
||||
}); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := queuepkg.EnqueueID(s.Queue, queuepkg.TypeAnalyzeEpisode, task.ID, 0); err != nil {
|
||||
if settleErr := s.failQueuedTask(task.ID, "queue_unavailable", "剧本反推任务入队失败"); settleErr != nil {
|
||||
return nil, fmt.Errorf("剧本反推任务入队失败且预扣返还失败: %w", settleErr)
|
||||
}
|
||||
_ = s.DB.Model(&model.CreativeProject{}).Where("id=?", projectID).Updates(map[string]any{"redraw_status": "review", "analysis_message": "剧本反推任务入队失败"}).Error
|
||||
return nil, errors.New("反推任务队列暂时不可用,请稍后重试")
|
||||
}
|
||||
return task, nil
|
||||
}
|
||||
|
||||
func (s *Creative) queueRedrawAnalysis(userID, projectID uuid.UUID, storyboardID *uuid.UUID) (*model.GenerationTask, error) {
|
||||
if s.Queue == nil {
|
||||
return nil, errors.New("反推任务队列不可用")
|
||||
}
|
||||
config, err := s.analysisModel(projectID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
requestPrefix := "analysis_"
|
||||
input := json.RawMessage(`{}`)
|
||||
if storyboardID != nil {
|
||||
requestPrefix = "storyboard_analysis_"
|
||||
input = json.RawMessage(`{"mode":"storyboard"}`)
|
||||
}
|
||||
task := &model.GenerationTask{
|
||||
RequestID: requestPrefix + uuid.NewString(), UserID: userID, ChannelID: &config.ChannelID,
|
||||
ModelID: &config.ModelID, ProjectID: &projectID, StoryboardID: storyboardID,
|
||||
TaskType: "prompt_reverse", Status: "submitted", InputData: input, EstimatedPoints: "0.00", PrepaidPoints: "0.00",
|
||||
}
|
||||
if err := s.DB.Transaction(func(tx *gorm.DB) error {
|
||||
var project model.CreativeProject
|
||||
if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).
|
||||
Where("id=? AND user_id=? AND project_type='video_redraw' AND deleted_at IS NULL", projectID, userID).
|
||||
Take(&project).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if project.SourceVideoAssetID == nil {
|
||||
return errors.New("请先上传原视频")
|
||||
}
|
||||
if project.AudioSource == "subtitle_file" && project.SubtitleAssetID == nil {
|
||||
return errors.New("当前选择字幕文件识别,请先上传字幕")
|
||||
}
|
||||
if storyboardID != nil {
|
||||
var storyboard model.EpisodeStoryboard
|
||||
if err := tx.Where("id=? AND project_id=? AND deleted_at IS NULL", *storyboardID, projectID).Take(&storyboard).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if storyboard.Locked {
|
||||
return errors.New("当前分镜已保护,无法重新反推")
|
||||
}
|
||||
}
|
||||
var active int64
|
||||
if err := tx.Model(&model.GenerationTask{}).
|
||||
Where("project_id=? AND task_type='prompt_reverse' AND status IN ?", projectID, activeTaskStatuses).
|
||||
Count(&active).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if active > 0 {
|
||||
return errors.New("当前剧本已有反推任务正在执行")
|
||||
}
|
||||
if err := billing.CreateTextGenerationTask(tx, task, config.Pricing, "视频反推"); err != nil {
|
||||
return err
|
||||
}
|
||||
if storyboardID == nil {
|
||||
return tx.Model(&model.CreativeProject{}).Where("id=?", projectID).
|
||||
Updates(map[string]any{"redraw_status": "analyzing", "analysis_message": "等待视频分析"}).Error
|
||||
}
|
||||
return nil
|
||||
}); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := queuepkg.EnqueueID(s.Queue, queuepkg.TypeAnalyzeEpisode, task.ID, 0); err != nil {
|
||||
if settleErr := s.failQueuedTask(task.ID, "queue_unavailable", "反推任务入队失败"); settleErr != nil {
|
||||
return nil, fmt.Errorf("反推任务入队失败且预扣返还失败: %w", settleErr)
|
||||
}
|
||||
return nil, errors.New("反推任务队列暂时不可用,请稍后重试")
|
||||
}
|
||||
return task, nil
|
||||
}
|
||||
|
||||
func (s *Creative) ResetRedrawAnalysis(userID, projectID uuid.UUID) ([]string, error) {
|
||||
media := newDeletionMediaSet()
|
||||
err := s.DB.Transaction(func(tx *gorm.DB) error {
|
||||
var project model.CreativeProject
|
||||
if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).
|
||||
Where("id=? AND user_id=? AND project_type='video_redraw' AND deleted_at IS NULL", projectID, userID).
|
||||
Take(&project).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
var active int64
|
||||
if err := tx.Model(&model.GenerationTask{}).Where("project_id=? AND status IN ?", projectID, activeTaskStatuses).Count(&active).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if active > 0 {
|
||||
return errors.New("当前剧本存在进行中的任务,暂时不能重新反推")
|
||||
}
|
||||
if err := removeRedrawDerivedContent(tx, projectID, media); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := media.deleteRows(tx); err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.Model(&project).Updates(map[string]any{"redraw_status": "uploaded", "analysis_message": nil, "redraw_script": nil}).Error
|
||||
})
|
||||
return media.objectKeys(), err
|
||||
}
|
||||
|
||||
func removeRedrawDerivedContent(tx *gorm.DB, projectID uuid.UUID, media *deletionMediaSet) error {
|
||||
taskScope := `project_id=? AND (task_type IN ('prompt_reverse','video_generation')
|
||||
OR (task_type='image_generation' AND input_data->>'target_type'='storyboard'))`
|
||||
if err := media.addQuery(tx.Table("media_assets media").Select("DISTINCT media.id,media.object_key").
|
||||
Joins("JOIN generation_outputs output ON output.media_asset_id=media.id").
|
||||
Joins("JOIN generation_tasks task ON task.id=output.task_id").Where(taskScope, projectID)); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := media.addQuery(tx.Table("media_assets media").Select("DISTINCT media.id,media.object_key").
|
||||
Joins("JOIN episode_storyboards storyboard ON storyboard.thumbnail_asset_id=media.id").
|
||||
Where("storyboard.project_id=?", projectID)); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := collectMediaByObjectKey(tx, media, fmt.Sprintf("juyou_ran/video-redraw/projects/%s/storyboards/%%", projectID)); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := collectMediaByObjectKey(tx, media, fmt.Sprintf("juyou_ran/video-redraw/projects/%s/episodes/%%/storyboards/%%", projectID)); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Model(&model.CreativeProject{}).Where("id=?", projectID).Update("cover_asset_id", nil).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Exec("UPDATE episode_storyboards SET active_output_id=NULL WHERE project_id=?", projectID).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Exec("DELETE FROM generation_outputs WHERE task_id IN (SELECT id FROM generation_tasks WHERE "+taskScope+")", projectID).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Exec("DELETE FROM generation_tasks WHERE "+taskScope, projectID).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.Exec("DELETE FROM episode_storyboards WHERE project_id=?", projectID).Error
|
||||
}
|
||||
@@ -0,0 +1,298 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"juhe-factory/api/internal/billing"
|
||||
"juhe-factory/api/internal/model"
|
||||
queuepkg "juhe-factory/api/internal/queue"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/clause"
|
||||
)
|
||||
|
||||
type ScriptAnalysisInput struct {
|
||||
Name string `json:"name"`
|
||||
SourceContent string `json:"source_content"`
|
||||
ResultContent string `json:"result_content"`
|
||||
}
|
||||
|
||||
type ScriptAnalysisImportProject struct {
|
||||
ID uuid.UUID `json:"id"`
|
||||
Name string `json:"name"`
|
||||
EpisodeCount int `json:"episode_count"`
|
||||
CharCount int `json:"char_count"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
}
|
||||
|
||||
func (s *Creative) ListScriptAnalyses(userID uuid.UUID) ([]model.ScriptAnalysis, error) {
|
||||
items := make([]model.ScriptAnalysis, 0)
|
||||
err := s.DB.Where("user_id=? AND deleted_at IS NULL", userID).Order("updated_at DESC").Find(&items).Error
|
||||
return items, err
|
||||
}
|
||||
|
||||
func (s *Creative) CreateScriptAnalysis(userID uuid.UUID) (*model.ScriptAnalysis, error) {
|
||||
item := &model.ScriptAnalysis{UserID: userID, Name: "未命名剧本", AnalysisResult: json.RawMessage(`{}`), AnalysisStatus: "idle"}
|
||||
if err := s.DB.Create(item).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return item, nil
|
||||
}
|
||||
|
||||
func (s *Creative) ListScriptAnalysisImportProjects(userID uuid.UUID) ([]ScriptAnalysisImportProject, error) {
|
||||
items := make([]ScriptAnalysisImportProject, 0)
|
||||
err := s.DB.Table("creative_projects project").
|
||||
Select(`project.id,project.name,project.updated_at,
|
||||
count(episode.id) FILTER (WHERE episode.deleted_at IS NULL AND episode.status='review' AND btrim(coalesce(episode.redraw_script,'')) <> '') AS episode_count,
|
||||
coalesce(sum(char_length(episode.redraw_script)) FILTER (WHERE episode.deleted_at IS NULL AND episode.status='review' AND btrim(coalesce(episode.redraw_script,'')) <> ''),0) AS char_count`).
|
||||
Joins("LEFT JOIN project_episodes episode ON episode.project_id=project.id").
|
||||
Where(`project.user_id=? AND project.project_type='video_redraw' AND project.deleted_at IS NULL AND
|
||||
EXISTS (
|
||||
SELECT 1 FROM project_episodes source_episode
|
||||
WHERE source_episode.project_id=project.id AND source_episode.deleted_at IS NULL
|
||||
AND source_episode.status='review' AND btrim(coalesce(source_episode.redraw_script,'')) <> ''
|
||||
)`, userID).
|
||||
Group("project.id").Order("project.updated_at DESC").Scan(&items).Error
|
||||
return items, err
|
||||
}
|
||||
|
||||
func (s *Creative) ImportScriptAnalysisProject(userID, analysisID, projectID uuid.UUID) (*model.ScriptAnalysis, error) {
|
||||
err := s.DB.Transaction(func(tx *gorm.DB) error {
|
||||
var project model.CreativeProject
|
||||
if err := tx.Select("id", "name").
|
||||
Where("id=? AND user_id=? AND project_type='video_redraw' AND deleted_at IS NULL", projectID, userID).
|
||||
Take(&project).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
var episodes []model.ProjectEpisode
|
||||
if err := tx.Select("episode_no", "redraw_script").
|
||||
Where("project_id=? AND deleted_at IS NULL AND status='review' AND btrim(coalesce(redraw_script,'')) <> ''", projectID).
|
||||
Order("episode_no").Find(&episodes).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if len(episodes) == 0 {
|
||||
return errors.New("该项目暂无可导入的剧本")
|
||||
}
|
||||
parts := make([]string, 0, len(episodes))
|
||||
for _, episode := range episodes {
|
||||
parts = append(parts, fmt.Sprintf("【第%d集】\n%s", episode.EpisodeNo, strings.TrimSpace(episode.RedrawScript)))
|
||||
}
|
||||
name := project.Name
|
||||
return replaceScriptAnalysisSource(tx, userID, analysisID, name, strings.Join(parts, "\n\n"))
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return s.GetScriptAnalysis(userID, analysisID)
|
||||
}
|
||||
|
||||
func (s *Creative) ImportScriptAnalysisFile(userID, analysisID uuid.UUID, filename string, data []byte) (*model.ScriptAnalysis, error) {
|
||||
extension := strings.ToLower(filepath.Ext(filename))
|
||||
if extension != ".txt" && extension != ".doc" {
|
||||
return nil, errors.New("仅支持TXT和DOC文件")
|
||||
}
|
||||
if len(data) > maxDramaImportBytes {
|
||||
return nil, errors.New("剧本文件不能超过3MB")
|
||||
}
|
||||
content, _, err := parseDramaSource(filename, data, "")
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
content = normalizeDramaText(content)
|
||||
if content == "" {
|
||||
return nil, errors.New("未读取到有效剧本正文")
|
||||
}
|
||||
name := strings.TrimSpace(strings.TrimSuffix(filepath.Base(filename), filepath.Ext(filename)))
|
||||
if name == "" {
|
||||
name = "未命名剧本"
|
||||
}
|
||||
if err := s.DB.Transaction(func(tx *gorm.DB) error {
|
||||
return replaceScriptAnalysisSource(tx, userID, analysisID, name, content)
|
||||
}); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return s.GetScriptAnalysis(userID, analysisID)
|
||||
}
|
||||
|
||||
func replaceScriptAnalysisSource(tx *gorm.DB, userID, analysisID uuid.UUID, name, content string) error {
|
||||
var item model.ScriptAnalysis
|
||||
if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).
|
||||
Where("id=? AND user_id=? AND deleted_at IS NULL", analysisID, userID).Take(&item).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if item.AnalysisStatus == "queued" || item.AnalysisStatus == "running" {
|
||||
return errors.New("剧本正在分析,暂时不能导入")
|
||||
}
|
||||
if err := tx.Where("script_analysis_id=?", analysisID).Delete(&model.ScriptAnalysisCharacter{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.Model(&item).Updates(map[string]any{
|
||||
"name": strings.TrimSpace(name),
|
||||
"source_content": content, "result_content": "",
|
||||
"analysis_result": json.RawMessage(`{}`), "analysis_status": "idle", "analysis_message": "",
|
||||
}).Error
|
||||
}
|
||||
|
||||
func (s *Creative) GetScriptAnalysis(userID, id uuid.UUID) (*model.ScriptAnalysis, error) {
|
||||
var item model.ScriptAnalysis
|
||||
if err := s.DB.Preload("Characters", func(db *gorm.DB) *gorm.DB { return db.Order("sort_order,name") }).
|
||||
Where("id=? AND user_id=? AND deleted_at IS NULL", id, userID).Take(&item).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &item, nil
|
||||
}
|
||||
|
||||
type scriptAnalysisModelConfig struct {
|
||||
ModelID uuid.UUID
|
||||
ChannelID uuid.UUID
|
||||
ModelName string
|
||||
BaseURL string
|
||||
APIKeyCiphertext string
|
||||
Pricing billing.TextPricing `gorm:"-"`
|
||||
}
|
||||
|
||||
func (s *Creative) scriptAnalysisModel(userID uuid.UUID) (scriptAnalysisModelConfig, error) {
|
||||
var config scriptAnalysisModelConfig
|
||||
result := s.DB.Raw(`SELECT model.id AS model_id,channel.id AS channel_id,model.name AS model_name,
|
||||
channel.base_url,channel.api_key_ciphertext
|
||||
FROM user_model_configs preference
|
||||
JOIN models model ON model.id=preference.model_id AND model.model_type='text' AND model.enabled=true AND model.deleted_at IS NULL
|
||||
JOIN channels channel ON channel.id=model.channel_id AND channel.enabled=true AND channel.deleted_at IS NULL
|
||||
WHERE preference.user_id=? AND preference.project_type='script_analysis' AND preference.model_type='text'`, userID).Scan(&config)
|
||||
if result.Error != nil {
|
||||
return config, result.Error
|
||||
}
|
||||
if result.RowsAffected == 0 {
|
||||
return config, errors.New("请先配置剧本分析文本模型")
|
||||
}
|
||||
pricing, err := billing.LoadTextPricing(s.DB, config.ModelID)
|
||||
if err != nil {
|
||||
return config, errors.New("剧本分析文本模型计费配置无效")
|
||||
}
|
||||
config.Pricing = pricing
|
||||
return config, nil
|
||||
}
|
||||
|
||||
func (s *Creative) QueueScriptAnalysis(userID, analysisID uuid.UUID) (*model.GenerationTask, error) {
|
||||
if s.Queue == nil {
|
||||
return nil, errors.New("剧本分析任务队列不可用")
|
||||
}
|
||||
config, err := s.scriptAnalysisModel(userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
task := &model.GenerationTask{
|
||||
RequestID: "script_analysis_" + uuid.NewString(), UserID: userID, ChannelID: &config.ChannelID,
|
||||
ModelID: &config.ModelID, ScriptAnalysisID: &analysisID, TaskType: "script_analysis", Status: "submitted",
|
||||
EstimatedPoints: "0.00", PrepaidPoints: "0.00",
|
||||
}
|
||||
err = s.DB.Transaction(func(tx *gorm.DB) error {
|
||||
var item model.ScriptAnalysis
|
||||
if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).
|
||||
Where("id=? AND user_id=? AND deleted_at IS NULL", analysisID, userID).Take(&item).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if strings.TrimSpace(item.SourceContent) == "" {
|
||||
return errors.New("请先填写剧本原文")
|
||||
}
|
||||
var active int64
|
||||
if err := tx.Model(&model.GenerationTask{}).
|
||||
Where("script_analysis_id=? AND status IN ?", analysisID, activeTaskStatuses).Count(&active).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if active > 0 {
|
||||
return errors.New("当前剧本正在分析,请勿重复提交")
|
||||
}
|
||||
prompt, err := selectedPromptContent(tx, userID, "剧本分析")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
task.PromptSnapshot = prompt
|
||||
task.InputData = mustJSON(map[string]string{"source_content": item.SourceContent})
|
||||
if err := billing.CreateTextGenerationTask(tx, task, config.Pricing, "剧本分析"); err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.Model(&item).Updates(map[string]any{"analysis_status": "queued", "analysis_message": "剧本分析排队中"}).Error
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := queuepkg.EnqueueID(s.Queue, queuepkg.TypeAnalyzeScript, task.ID, 0); err != nil {
|
||||
if settleErr := s.failQueuedTask(task.ID, "queue_unavailable", "剧本分析任务入队失败"); settleErr != nil {
|
||||
return nil, fmt.Errorf("剧本分析任务入队失败且预扣返还失败: %w", settleErr)
|
||||
}
|
||||
_ = s.DB.Model(&model.ScriptAnalysis{}).Where("id=?", analysisID).
|
||||
Updates(map[string]any{"analysis_status": "failed", "analysis_message": "剧本分析任务入队失败"}).Error
|
||||
return nil, errors.New("剧本分析任务队列暂时不可用,请稍后重试")
|
||||
}
|
||||
return task, nil
|
||||
}
|
||||
|
||||
func (s *Creative) CancelScriptAnalysis(userID, analysisID uuid.UUID) error {
|
||||
return s.DB.Transaction(func(tx *gorm.DB) error {
|
||||
var task model.GenerationTask
|
||||
if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).
|
||||
Where("script_analysis_id=? AND user_id=? AND task_type='script_analysis' AND status IN ?", analysisID, userID,
|
||||
[]string{"submitted", "processing", "cancel_requested"}).
|
||||
Order("created_at DESC").Take(&task).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if task.Status == "cancel_requested" {
|
||||
return nil
|
||||
}
|
||||
if task.Status == "submitted" {
|
||||
if err := tx.Model(&task).Updates(map[string]any{
|
||||
"status": "cancelled", "actual_points": task.PrepaidPoints, "finished_at": time.Now(),
|
||||
"error_code": "user_cancelled", "error_message": "用户取消,已扣积分不退",
|
||||
}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.Model(&model.ScriptAnalysis{}).Where("id=? AND user_id=? AND deleted_at IS NULL", analysisID, userID).Updates(map[string]any{
|
||||
"analysis_status": "idle", "analysis_message": "剧本分析已取消,已扣积分不退",
|
||||
}).Error
|
||||
}
|
||||
if err := tx.Model(&task).Updates(map[string]any{
|
||||
"status": "cancel_requested", "error_code": "user_cancelled", "error_message": "正在取消分析",
|
||||
}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.Model(&model.ScriptAnalysis{}).Where("id=? AND user_id=? AND deleted_at IS NULL", analysisID, userID).
|
||||
Update("analysis_message", "正在取消分析").Error
|
||||
})
|
||||
}
|
||||
|
||||
func (s *Creative) UpdateScriptAnalysis(userID, id uuid.UUID, input ScriptAnalysisInput) (*model.ScriptAnalysis, error) {
|
||||
name := strings.TrimSpace(input.Name)
|
||||
if name == "" {
|
||||
return nil, errors.New("剧本名称不能为空")
|
||||
}
|
||||
result := s.DB.Model(&model.ScriptAnalysis{}).
|
||||
Where("id=? AND user_id=? AND deleted_at IS NULL", id, userID).
|
||||
Updates(map[string]any{"name": name, "source_content": input.SourceContent, "result_content": input.ResultContent})
|
||||
if result.Error != nil {
|
||||
return nil, result.Error
|
||||
}
|
||||
if result.RowsAffected == 0 {
|
||||
return nil, gorm.ErrRecordNotFound
|
||||
}
|
||||
return s.GetScriptAnalysis(userID, id)
|
||||
}
|
||||
|
||||
func (s *Creative) DeleteScriptAnalysis(userID, id uuid.UUID) error {
|
||||
result := s.DB.Model(&model.ScriptAnalysis{}).
|
||||
Where("id=? AND user_id=? AND deleted_at IS NULL", id, userID).
|
||||
Update("deleted_at", gorm.Expr("CURRENT_TIMESTAMP"))
|
||||
if result.Error != nil {
|
||||
return result.Error
|
||||
}
|
||||
if result.RowsAffected == 0 {
|
||||
return gorm.ErrRecordNotFound
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,24 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/google/uuid"
|
||||
)
|
||||
|
||||
func TestImportScriptAnalysisFileRejectsUnsupportedFormat(t *testing.T) {
|
||||
service := &Creative{}
|
||||
_, err := service.ImportScriptAnalysisFile(uuid.New(), uuid.New(), "script.docx", []byte("content"))
|
||||
if err == nil || !strings.Contains(err.Error(), "仅支持TXT和DOC文件") {
|
||||
t.Fatalf("expected unsupported format error, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestImportScriptAnalysisFileRejectsOversizedFile(t *testing.T) {
|
||||
service := &Creative{}
|
||||
_, err := service.ImportScriptAnalysisFile(uuid.New(), uuid.New(), "script.txt", make([]byte, maxDramaImportBytes+1))
|
||||
if err == nil || !strings.Contains(err.Error(), "不能超过3MB") {
|
||||
t.Fatalf("expected file size error, got %v", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,272 @@
|
||||
// WEB 用户服务,负责用户认证、账户资料、积分查询和兑换码业务。
|
||||
package service
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"strings"
|
||||
"time"
|
||||
"unicode/utf8"
|
||||
|
||||
"juhe-factory/api/internal/billing"
|
||||
"juhe-factory/api/internal/model"
|
||||
"juhe-factory/api/internal/security"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/clause"
|
||||
)
|
||||
|
||||
type Web struct {
|
||||
DB *gorm.DB
|
||||
Passwords security.PasswordHasher
|
||||
Tokens security.WebTokenService
|
||||
RefreshTTL time.Duration
|
||||
}
|
||||
|
||||
var (
|
||||
ErrInvalidUsername = errors.New("用户名长度必须为 5~12 个字符")
|
||||
ErrUsernameExists = errors.New("用户名已被使用")
|
||||
ErrOriginalPassword = errors.New("原密码错误")
|
||||
ErrAccountDisabled = errors.New("账户已被禁用")
|
||||
)
|
||||
|
||||
type WebTokenPair struct {
|
||||
AccessToken string `json:"access_token"`
|
||||
RefreshToken string `json:"refresh_token"`
|
||||
ExpiresAt time.Time `json:"expires_at"`
|
||||
User map[string]any `json:"user"`
|
||||
}
|
||||
|
||||
func (s *Web) Login(username, password string) (*WebTokenPair, error) {
|
||||
var user model.WebUser
|
||||
if err := s.DB.Where("account = ? AND deleted_at IS NULL", strings.TrimSpace(username)).First(&user).Error; err != nil {
|
||||
return nil, ErrUnauthorized
|
||||
}
|
||||
if !user.Enabled {
|
||||
return nil, ErrAccountDisabled
|
||||
}
|
||||
if !s.Passwords.Verify(user.PasswordHash, password) {
|
||||
return nil, ErrUnauthorized
|
||||
}
|
||||
if user.LastOnlineAt == nil || user.LastOnlineAt.Before(time.Now().Add(-5*time.Minute)) {
|
||||
_ = s.DB.Model(&user).Update("last_online_at", time.Now()).Error
|
||||
}
|
||||
return s.issue(&user)
|
||||
}
|
||||
|
||||
func (s *Web) Refresh(refreshToken string) (*WebTokenPair, error) {
|
||||
var record model.WebRefreshToken
|
||||
if err := s.DB.Where("token_hash = ? AND revoked_at IS NULL AND expires_at > ?", security.HashToken(refreshToken), time.Now()).First(&record).Error; err != nil {
|
||||
return nil, ErrUnauthorized
|
||||
}
|
||||
var user model.WebUser
|
||||
if err := s.DB.Where("id = ? AND enabled = true AND deleted_at IS NULL", record.UserID).First(&user).Error; err != nil || user.SessionVersion != record.SessionVersion {
|
||||
return nil, ErrUnauthorized
|
||||
}
|
||||
now := time.Now()
|
||||
if err := s.DB.Model(&record).Update("revoked_at", now).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return s.issue(&user)
|
||||
}
|
||||
|
||||
func (s *Web) Logout(refreshToken string) error {
|
||||
if refreshToken == "" {
|
||||
return nil
|
||||
}
|
||||
return s.DB.Model(&model.WebRefreshToken{}).Where("token_hash = ? AND revoked_at IS NULL", security.HashToken(refreshToken)).Update("revoked_at", time.Now()).Error
|
||||
}
|
||||
|
||||
func (s *Web) Authenticate(raw string) (*model.WebUser, error) {
|
||||
claims, err := s.Tokens.ParseAccessToken(raw)
|
||||
if err != nil {
|
||||
return nil, ErrUnauthorized
|
||||
}
|
||||
id, err := uuid.Parse(claims.UserID)
|
||||
if err != nil {
|
||||
return nil, ErrUnauthorized
|
||||
}
|
||||
var user model.WebUser
|
||||
if err := s.DB.Where("id = ? AND enabled = true AND deleted_at IS NULL", id).First(&user).Error; err != nil || user.SessionVersion != claims.SessionVersion {
|
||||
return nil, ErrUnauthorized
|
||||
}
|
||||
return &user, nil
|
||||
}
|
||||
|
||||
func (s *Web) issue(user *model.WebUser) (*WebTokenPair, error) {
|
||||
access, expires, err := s.Tokens.NewAccessToken(user.ID.String(), user.UID, user.Account, user.SessionVersion)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
opaque := security.TokenService{RefreshTTL: s.RefreshTTL}
|
||||
refresh, hash, refreshExpires, err := opaque.NewRefreshToken()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
record := model.WebRefreshToken{UserID: user.ID, TokenHash: hash, SessionVersion: user.SessionVersion, ExpiresAt: refreshExpires}
|
||||
if err := s.DB.Create(&record).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &WebTokenPair{AccessToken: access, RefreshToken: refresh, ExpiresAt: expires, User: map[string]any{"id": user.ID, "uid": user.UID, "username": user.Username, "account": user.Account, "avatar_url": user.AvatarURL}}, nil
|
||||
}
|
||||
|
||||
// Account 返回用户资料、永久积分余额、每日限额和当日消耗数据。
|
||||
func (s *Web) Account(userID uuid.UUID) (map[string]any, error) {
|
||||
result := map[string]any{}
|
||||
err := s.DB.Raw(`SELECT id,uid,username,account,avatar_url,point_balance,daily_limit,last_online_at,
|
||||
coalesce((SELECT sum(-hold.change_amount) FROM point_ledger hold
|
||||
WHERE hold.user_id=web_users.id AND hold.business_type='generation_hold'
|
||||
AND hold.created_at>=CURRENT_DATE AND hold.created_at<CURRENT_DATE+INTERVAL '1 day'
|
||||
AND NOT EXISTS (SELECT 1 FROM point_ledger refund WHERE refund.idempotency_key='generation:refund:'||hold.business_id)),0) AS today_consumption
|
||||
FROM web_users WHERE id=? AND deleted_at IS NULL`, userID).Scan(&result).Error
|
||||
return result, err
|
||||
}
|
||||
|
||||
func (s *Web) UpdateProfile(userID uuid.UUID, username string) (map[string]any, error) {
|
||||
username = strings.TrimSpace(username)
|
||||
if count := utf8.RuneCountInString(username); count < 5 || count > 12 {
|
||||
return nil, ErrInvalidUsername
|
||||
}
|
||||
var duplicates int64
|
||||
if err := s.DB.Table("web_users").Where("username = ? AND id <> ?", username, userID).Count(&duplicates).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if duplicates > 0 {
|
||||
return nil, ErrUsernameExists
|
||||
}
|
||||
if err := s.DB.Table("web_users").Where("id = ? AND deleted_at IS NULL", userID).Update("username", username).Error; err != nil {
|
||||
var raced int64
|
||||
if s.DB.Table("web_users").Where("username = ? AND id <> ?", username, userID).Count(&raced).Error == nil && raced > 0 {
|
||||
return nil, ErrUsernameExists
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return s.Account(userID)
|
||||
}
|
||||
|
||||
func (s *Web) ChangePassword(userID uuid.UUID, originalPassword, newPassword string) error {
|
||||
if len(newPassword) < 8 || len(newPassword) > 128 {
|
||||
return errors.New("新密码长度必须为 8~128 位")
|
||||
}
|
||||
var user model.WebUser
|
||||
if err := s.DB.Where("id = ? AND enabled = true AND deleted_at IS NULL", userID).First(&user).Error; err != nil {
|
||||
return ErrUnauthorized
|
||||
}
|
||||
if !s.Passwords.Verify(user.PasswordHash, originalPassword) {
|
||||
return ErrOriginalPassword
|
||||
}
|
||||
hash, err := s.Passwords.Hash(newPassword)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return s.DB.Model(&user).Update("password_hash", hash).Error
|
||||
}
|
||||
|
||||
func (s *Web) AvatarKey(userID uuid.UUID) (string, error) {
|
||||
var key string
|
||||
err := s.DB.Raw("SELECT coalesce(avatar_key, '') FROM web_users WHERE id = ? AND deleted_at IS NULL", userID).Scan(&key).Error
|
||||
return key, err
|
||||
}
|
||||
|
||||
func (s *Web) UpdateAvatar(userID uuid.UUID, url, key string) error {
|
||||
return s.DB.Table("web_users").Where("id = ? AND deleted_at IS NULL", userID).Updates(map[string]any{"avatar_url": url, "avatar_key": key}).Error
|
||||
}
|
||||
|
||||
func (s *Web) Usage30Days(userID uuid.UUID) ([]map[string]any, error) {
|
||||
items := make([]map[string]any, 0)
|
||||
err := s.DB.Raw(`WITH days AS (SELECT generate_series((CURRENT_DATE-INTERVAL '29 days')::date,CURRENT_DATE::date,'1 day')::date AS day),
|
||||
stats AS (SELECT hold.created_at::date AS day,coalesce(sum(-hold.change_amount),0) AS consumption
|
||||
FROM point_ledger hold WHERE hold.user_id=? AND hold.business_type='generation_hold'
|
||||
AND hold.created_at>=CURRENT_DATE-INTERVAL '29 days'
|
||||
AND NOT EXISTS (SELECT 1 FROM point_ledger refund WHERE refund.idempotency_key='generation:refund:'||hold.business_id)
|
||||
GROUP BY hold.created_at::date)
|
||||
SELECT days.day,coalesce(stats.consumption,0) AS consumption FROM days LEFT JOIN stats USING(day) ORDER BY days.day`, userID).Scan(&items).Error
|
||||
return items, err
|
||||
}
|
||||
|
||||
// ConsumptionRecordPage 描述用户积分流水的服务端分页结果。
|
||||
type ConsumptionRecordPage struct {
|
||||
Items []map[string]any `json:"items"`
|
||||
Total int64 `json:"total"`
|
||||
Page int `json:"page"`
|
||||
PageSize int `json:"page_size"`
|
||||
}
|
||||
|
||||
// normalizeConsumptionPagination 约束积分流水页码和每页数量,避免无效或过大的查询。
|
||||
func normalizeConsumptionPagination(page, pageSize int) (int, int) {
|
||||
if page < 1 {
|
||||
page = 1
|
||||
}
|
||||
if pageSize < 1 {
|
||||
pageSize = 10
|
||||
}
|
||||
if pageSize > 100 {
|
||||
pageSize = 100
|
||||
}
|
||||
return page, pageSize
|
||||
}
|
||||
|
||||
// ConsumptionRecords 返回用户消费、退款、兑换及后台发放形成的积分变动分页流水。
|
||||
func (s *Web) ConsumptionRecords(userID uuid.UUID, page, pageSize int) (ConsumptionRecordPage, error) {
|
||||
page, pageSize = normalizeConsumptionPagination(page, pageSize)
|
||||
result := ConsumptionRecordPage{Items: make([]map[string]any, 0), Page: page, PageSize: pageSize}
|
||||
ledger := s.DB.Table("point_ledger").
|
||||
Where("user_id=? AND business_type IN ?", userID, []string{"generation_hold", "generation_refund", "redemption", "admin_grant"})
|
||||
if err := ledger.Count(&result.Total).Error; err != nil {
|
||||
return result, err
|
||||
}
|
||||
offset := (page - 1) * pageSize
|
||||
err := s.DB.Raw(`SELECT id,abs(change_amount)::text AS amount,change_amount::text AS change_amount,
|
||||
balance_after::text AS balance_after,business_type,remark,created_at
|
||||
FROM point_ledger
|
||||
WHERE user_id=? AND business_type IN ('generation_hold','generation_refund','redemption','admin_grant')
|
||||
ORDER BY created_at DESC,id DESC LIMIT ? OFFSET ?`, userID, pageSize, offset).Scan(&result.Items).Error
|
||||
return result, err
|
||||
}
|
||||
|
||||
// Redeem 校验兑换码并将对应积分以永久有效批次计入用户账户。
|
||||
func (s *Web) Redeem(userID uuid.UUID, code string) (map[string]any, error) {
|
||||
normalized := strings.ReplaceAll(strings.TrimSpace(code), "-", "")
|
||||
if normalized == "" {
|
||||
return nil, errors.New("请输入兑换码")
|
||||
}
|
||||
sum := sha256.Sum256([]byte(normalized))
|
||||
hash := hex.EncodeToString(sum[:])
|
||||
result := map[string]any{}
|
||||
err := s.DB.Transaction(func(tx *gorm.DB) error {
|
||||
var record struct {
|
||||
ID uuid.UUID
|
||||
Points string
|
||||
Status string
|
||||
ExpiresAt *time.Time
|
||||
}
|
||||
if err := tx.Table("redemption_codes").Clauses(clause.Locking{Strength: "UPDATE"}).Select("id,points::text AS points,status,expires_at").Where("code_hash=?", hash).Take(&record).Error; err != nil {
|
||||
return errors.New("兑换码不存在")
|
||||
}
|
||||
if record.Status != "unused" {
|
||||
return errors.New("兑换码已使用或已过期")
|
||||
}
|
||||
if record.ExpiresAt != nil && !record.ExpiresAt.After(time.Now()) {
|
||||
if err := tx.Table("redemption_codes").Where("id=?", record.ID).Update("status", "expired").Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return errors.New("兑换码已过期")
|
||||
}
|
||||
redeemedAt := time.Now()
|
||||
balance, credited, err := billing.CreditRedemptionPoints(tx, userID, record.ID, record.Points)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !credited {
|
||||
return errors.New("兑换码积分已入账")
|
||||
}
|
||||
if err := tx.Table("redemption_codes").Where("id=?", record.ID).Updates(map[string]any{"status": "used", "redeemed_by": userID, "redeemed_at": redeemedAt}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
result = map[string]any{"points": record.Points, "balance": balance}
|
||||
return nil
|
||||
})
|
||||
return result, err
|
||||
}
|
||||
@@ -0,0 +1,27 @@
|
||||
// 用户积分流水分页测试,验证无效参数修正和每页数量上限。
|
||||
package service
|
||||
|
||||
import "testing"
|
||||
|
||||
// TestNormalizeConsumptionPagination 验证积分流水分页参数始终落在允许范围内。
|
||||
func TestNormalizeConsumptionPagination(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
page int
|
||||
pageSize int
|
||||
wantPage int
|
||||
wantPageSize int
|
||||
}{
|
||||
{name: "defaults", page: 0, pageSize: 0, wantPage: 1, wantPageSize: 10},
|
||||
{name: "keeps valid values", page: 3, pageSize: 20, wantPage: 3, wantPageSize: 20},
|
||||
{name: "caps page size", page: 2, pageSize: 500, wantPage: 2, wantPageSize: 100},
|
||||
}
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
page, pageSize := normalizeConsumptionPagination(test.page, test.pageSize)
|
||||
if page != test.wantPage || pageSize != test.wantPageSize {
|
||||
t.Fatalf("pagination = (%d, %d), want (%d, %d)", page, pageSize, test.wantPage, test.wantPageSize)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user