diff --git a/Dockerfile b/Dockerfile index 561f976..30b64a0 100644 --- a/Dockerfile +++ b/Dockerfile @@ -13,12 +13,13 @@ RUN go build -ldflags="-s -w" -o main ./main.go FROM alpine:3.19 RUN sed -i 's/dl-cdn.alpinelinux.org/mirrors.aliyun.com/g' /etc/apk/repositories \ - && apk add --no-cache ca-certificates tzdata + && apk add --no-cache ca-certificates tzdata ffmpeg ENV TZ=Asia/Shanghai RUN ln -snf /usr/share/zoneinfo/$TZ /etc/localtime && echo $TZ > /etc/timezone WORKDIR /app COPY --from=builder /build/config.yml . COPY --from=builder /build/prompt.md . +COPY --from=builder /build/negative_prompt.md . COPY --from=builder /build/default_first_frame.png . COPY --from=builder /build/main . RUN printf '#!/bin/sh\nif [ -d /app/short_drama.db ]; then rm -rf /app/short_drama.db; fi\ntouch /app/short_drama.db 2>/dev/null || true\nexec ./main\n' > /app/entrypoint.sh && chmod +x /app/entrypoint.sh diff --git a/api.json b/api.json new file mode 100644 index 0000000..a44d3f0 --- /dev/null +++ b/api.json @@ -0,0 +1,218 @@ +{ + "model": "wan2.6-r2v-flash", + "input": { + "prompt": "Character2 坐在靠窗的椅子上,手持 character3,在 character4 旁演奏一首舒缓的美国乡村民谣。Character1 对Character2开口说道:“听起来不错”", + "negative_prompt": "", + "reference_urls": [ + "https://help-static-aliyun-doc.aliyuncs.com/file-manage-files/zh-CN/20260129/hfugmr/wan-r2v-role1.mp4", + "https://help-static-aliyun-doc.aliyuncs.com/file-manage-files/zh-CN/20260129/qigswt/wan-r2v-role2.mp4", + "https://help-static-aliyun-doc.aliyuncs.com/file-manage-files/zh-CN/20260129/qpzxps/wan-r2v-object4.png", + "https://help-static-aliyun-doc.aliyuncs.com/file-manage-files/zh-CN/20260129/wfjikw/wan-r2v-backgroud5.png" + ] + }, + "parameters": { + "size": "1280*720", + "duration": 10, + "audio": true, + "shot_type": "multi", + "watermark": true, + "seed": 2147483647 + } +} + +//model string (必选) +// +//模型名称。模型列表与价格详见模型价格。 +// +//示例值:wan2.6-r2v-flash。 +// +//input object (必选) +// +//输入的基本信息,如提示词等。 +// +//属性 +// +//prompt string (必选) +// +//文本提示词。用来描述生成视频中期望包含的元素和视觉特点。 +// +//支持中英文,每个汉字、字母、标点占一个字符,超过部分会自动截断。 +// +//wan2.6-r2v-flash:长度不超过1500个字符。 +// +//wan2.6-r2v:长度不超过1500个字符。 +// +//角色引用说明:通过“character1、character2”这类标识引用参考角色,每个参考(视频或图像)仅包含单一角色。模型仅通过此方式识别参考中的角色。 +// +//示例值:character1在沙发上开心地看电影。 +// +//提示词的使用技巧请参见文生视频/图生视频Prompt指南。 +// +//negative_prompt string (可选) +// +//反向提示词,用来描述不希望在视频画面中出现的内容,可以对视频画面进行限制。 +// +//支持中英文,长度不超过500个字符,超过部分会自动截断。 +// +//示例值:低分辨率、错误、最差质量、低质量、残缺、多余的手指、比例不良等。 +// +//reference_urls array[string] (必选) +// +//重要 +//reference_urls直接影响费用,计费规则请参见计费与限流。 +// +//上传的参考文件 URL 数组,支持传入视频和图像。用于提取角色形象与音色(如有),以生成符合参考特征的视频。 +// +//每个 URL 可指向 一张图像 或 一段视频: +// +//图像数量:0~5。 +// +//视频数量:0~3。 +// +//总数限制:图像 + 视频 ≤ 5。 +// +//传入多个参考文件时,按照数组顺序定义角色的顺序。即第 1 个 URL 对应 character1,第 2 个对应 character2,以此类推。 +// +//每个参考文件仅包含一个主体角色。例如 character1 为小女孩,character2 为闹钟。 +// +//支持输入的格式: +// +//公网URL: +// +//支持 HTTP 或 HTTPS 协议。 +// +//示例值:https://cdn.translate.alibaba.com/xxx.png。 +// +//临时URL: +// +//支持OSS协议,必须通过上传文件获取临时 URL。 +// +//示例值:oss://dashscope-instant/xxx/xxx.png。 +// +//参考视频要求: +// +//格式:MP4、MOV。 +// +//时长:1s~30s。 +// +//视频大小:不超过100MB。 +// +//参考图像要求: +// +//格式:JPEG、JPG、PNG(不支持透明通道)、BMP、WEBP。 +// +//分辨率:宽高均需在[240,8000]像素之间。 +// +//图像大小:不超过20MB。 +// +//示例值:["https://help-static-aliyun-doc.aliyuncs.com/file-manage-files/xxx.mp4", "https://help-static-aliyun-doc.aliyuncs.com/file-manage-files/xxx.jpg"]。 +// +//已废弃字段 +// +//parameters object (可选) +// +//图像处理参数。如设置视频分辨率、开启prompt智能改写、添加水印等。 +// +//属性 +// +//size string (可选) +// +//重要 +//size直接影响费用,费用 = 单价(基于分辨率)× 时长(秒)。同一模型:1080P > 720P ,请在调用前确认模型价格。 +// +//size必须设置为具体数值(如 1280*720),而不是 1:1或720P。 +// +//指定生成的视频分辨率,格式为宽*高。该参数的默认值和可用枚举值依赖于 model 参数,规则如下: +// +//wan2.6-r2v-flash:默认值为 1920*1080(1080P)。可选分辨率:720P、1080P对应的所有分辨率。 +// +//wan2.6-r2v:默认值为 1920*1080(1080P)。可选分辨率:720P、1080P对应的所有分辨率。 +// +//720P档位:可选的视频分辨率及其对应的视频宽高比为: +// +//1280*720:16:9。 +// +//720*1280:9:16。 +// +//960*960:1:1。 +// +//1088*832:4:3。 +// +//832*1088:3:4。 +// +//1080P档位:可选的视频分辨率及其对应的视频宽高比为: +// +//1920*1080: 16:9。 +// +//1080*1920: 9:16。 +// +//1440*1440: 1:1。 +// +//1632*1248: 4:3。 +// +//1248*1632: 3:4。 +// +//duration integer (可选) +// +//重要 +//duration直接影响费用。费用 = 单价(基于分辨率)× 时长(秒),请在调用前确认模型价格。 +// +//生成视频的时长,单位为秒。 +// +//wan2.6-r2v-flash:取值为[2, 10]之间的整数。默认值为5。 +// +//wan2.6-r2v:取值为[2, 10]之间的整数。默认值为5。 +// +//示例值:5。 +// +//shot_type string (可选) +// +//指定生成视频的镜头类型,即视频是由一个连续镜头还是多个切换镜头组成。 +// +//参数优先级:shot_type > prompt。例如,若 shot_type设置为"single",即使 prompt 中包含“生成多镜头视频”,模型仍会输出单镜头视频。 +// +//可选值: +// +//single:默认值,输出单镜头视频 +// +//multi:输出多镜头视频。 +// +//示例值:single。 +// +//说明 +//当希望严格控制视频的叙事结构(如产品展示用单镜头、故事短片用多镜头),可通过此参数指定。 +// +//audio boolean (可选) +// +//重要 +//audio直接影响费用,有声视频与无声视频价格不同,请在调用前确认模型价格。 +// +//支持模型:wan2.6-r2v-flash。 +// +//是否生成有声视频。 +// +//可选值: +// +//true:默认值,输出有声视频。 +// +//false:输出无声视频。 +// +//示例值:true。 +// +//watermark boolean (可选) +// +//是否添加水印标识,水印位于视频右下角,文案固定为“AI生成”。 +// +//false:默认值,不添加水印。 +// +//true:添加水印。 +// +//示例值:false。 +// +//seed integer (可选) +// +//随机数种子,取值范围为[0, 2147483647]。 +// +//未指定时,系统自动生成随机种子。若需提升生成结果的可复现性,建议固定seed值。 +// +//请注意,由于模型生成具有概率性,即使使用相同 seed,也不能保证每次生成结果完全一致。 \ No newline at end of file diff --git a/default_first_frame.png b/default_first_frame.png index 94381b4..9fce5c0 100644 Binary files a/default_first_frame.png and b/default_first_frame.png differ diff --git a/negative_prompt.md b/negative_prompt.md new file mode 100644 index 0000000..91d8559 --- /dev/null +++ b/negative_prompt.md @@ -0,0 +1 @@ +低分辨率、模糊、失真、扭曲、变形、闪烁、抖动、过度曝光、色彩失真、画面杂乱、构图不当、主体不完整、多余物体、残缺、面部扭曲、肢体不自然、比例失调、低质量、最差质量、不良渲染 diff --git a/short_drama.db b/short_drama.db index aea8100..6503b22 100644 Binary files a/short_drama.db and b/short_drama.db differ diff --git a/shortdrama/agent/chat_model.go b/shortdrama/agent/chat_model.go index 16ea6f6..463c402 100644 --- a/shortdrama/agent/chat_model.go +++ b/shortdrama/agent/chat_model.go @@ -6,10 +6,11 @@ import ( "encoding/json" "fmt" "io" - "log" "net/http" "strings" "time" + + "github.com/gogf/gf/v2/frame/g" ) // ModelConfig 模型配置 @@ -39,7 +40,7 @@ func CallChatModel(ctx context.Context, cfg *ModelConfig, req *ChatRequest) (*Ch timeout := cfg.Timeout if timeout <= 0 { - timeout = 60 * time.Second + timeout = 180 * time.Second } body, err := buildReqBody(cfg.ModelName, req) @@ -48,13 +49,14 @@ func CallChatModel(ctx context.Context, cfg *ModelConfig, req *ChatRequest) (*Ch } url := trimSlashes(cfg.BaseURL) + "/v1/chat/completions" + g.Log().Infof(ctx, "ChatAPI 开始调用 model=%s timeout=%v max_retries=3 body_size=%d", cfg.ModelName, timeout, len(body)) var lastErr error maxRetries := 3 for attempt := 0; attempt <= maxRetries; attempt++ { if attempt > 0 { wait := time.Duration(1<<(attempt-1)) * time.Second - log.Printf("API限流重试第%d次(等待%v)...", attempt, wait) + g.Log().Infof(ctx, "ChatAPI 限流重试第%d次(等待%v)", attempt, wait) select { case <-ctx.Done(): return nil, ctx.Err() @@ -64,10 +66,13 @@ func CallChatModel(ctx context.Context, cfg *ModelConfig, req *ChatRequest) (*Ch result, doErr := doChatRequest(ctx, url, cfg.APIKey, body, timeout) if doErr == nil { + g.Log().Infof(ctx, "ChatAPI 调用成功 url=%s tool_calls=%d content_len=%d", + url, len(result.ToolCalls), len(result.Content)) return result, nil } lastErr = doErr + g.Log().Warningf(ctx, "ChatAPI 请求失败(attempt=%d/%d): %v", attempt+1, maxRetries+1, doErr) // 只有限流或服务端错误才重试 errStr := lastErr.Error() if !strings.Contains(errStr, "limit_requests") && @@ -79,6 +84,7 @@ func CallChatModel(ctx context.Context, cfg *ModelConfig, req *ChatRequest) (*Ch } } + g.Log().Errorf(ctx, "ChatAPI %d次重试后最终失败: %v", maxRetries+1, lastErr) return nil, lastErr } @@ -90,18 +96,27 @@ func doChatRequest(ctx context.Context, url, apiKey string, body []byte, timeout httpReq.Header.Set("Authorization", "Bearer "+apiKey) httpReq.Header.Set("Content-Type", "application/json") + start := time.Now() client := &http.Client{Timeout: timeout} resp, err := client.Do(httpReq) + elapsed := time.Since(start) if err != nil { - return nil, fmt.Errorf("请求失败: %w", err) + return nil, fmt.Errorf("请求失败(耗时%v): %w", elapsed, err) } defer resp.Body.Close() respBody, err := io.ReadAll(resp.Body) if err != nil { - return nil, fmt.Errorf("读取响应失败: %w", err) + return nil, fmt.Errorf("读取响应失败(状态码=%d): %w", resp.StatusCode, err) } + if resp.StatusCode != 200 { + return nil, fmt.Errorf("API返回错误状态码=%d body=%s", resp.StatusCode, string(respBody)) + } + + g.Log().Infof(ctx, "ChatAPI 响应完成 status=%d body_len=%d elapsed=%v", + resp.StatusCode, len(respBody), elapsed) + return parseRespBody(respBody) } @@ -145,24 +160,37 @@ type apiRespBody struct { } type apiChoice struct { - Index int `json:"index"` - Message apiMsg `json:"message"` + Index int `json:"index"` + Message apiRespMsg `json:"message"` } -type apiMsg struct { - Content string `json:"content"` - ToolCalls []apiToolCall `json:"tool_calls,omitempty"` +// apiRespMsg 响应消息体(arguments 使用 json.RawMessage 兼容对象和字符串) +type apiRespMsg struct { + Content string `json:"content"` + ToolCalls []apiRespToolCall `json:"tool_calls,omitempty"` +} + +type apiRespToolCall struct { + ID string `json:"id"` + Type string `json:"type"` + Function apiRespFuncCall `json:"function"` +} + +type apiRespFuncCall struct { + Name string `json:"name"` + Arguments json.RawMessage `json:"arguments"` } type apiToolCall struct { - ID string `json:"id"` - Type string `json:"type"` - Function apiFuncCall `json:"function"` + ID string `json:"id"` + Type string `json:"type"` + Function apiReqFuncCall `json:"function"` } -type apiFuncCall struct { - Name string `json:"name"` - Arguments string `json:"arguments"` +// apiReqFuncCall 请求中的 function call(arguments 为 json.RawMessage 避免二次编码) +type apiReqFuncCall struct { + Name string `json:"name"` + Arguments json.RawMessage `json:"arguments"` } func buildReqBody(model string, req *ChatRequest) ([]byte, error) { @@ -201,12 +229,16 @@ func toAPIMessages(msgs []*ChatMessage) []apiMessage { if len(m.ToolCalls) > 0 { om.ToolCalls = make([]apiToolCall, 0, len(m.ToolCalls)) for _, tc := range m.ToolCalls { + args := tc.Arguments + if args == "" { + args = "{}" + } om.ToolCalls = append(om.ToolCalls, apiToolCall{ ID: tc.ID, Type: "function", - Function: apiFuncCall{ + Function: apiReqFuncCall{ Name: tc.Name, - Arguments: tc.Arguments, + Arguments: json.RawMessage(args), }, }) } @@ -233,16 +265,34 @@ func parseRespBody(data []byte) (*ChatResponse, error) { if len(msg.ToolCalls) > 0 { cr.ToolCalls = make([]*ToolCall, 0, len(msg.ToolCalls)) for _, tc := range msg.ToolCalls { + args := resolveArguments(tc.Function.Arguments) cr.ToolCalls = append(cr.ToolCalls, &ToolCall{ ID: tc.ID, Name: tc.Function.Name, - Arguments: tc.Function.Arguments, + Arguments: args, }) } } return cr, nil } +// resolveArguments 将 json.RawMessage 的参数转为字符串 +// API 可能返回 "arguments": "{\"key\":\"val\"}"(字符串)或 "arguments": {"key":"val"}(对象) +func resolveArguments(raw json.RawMessage) string { + if len(raw) == 0 { + return "" + } + // 如果是 JSON 字符串(以 " 开头),直接提取字符串值 + if raw[0] == '"' { + var s string + if json.Unmarshal(raw, &s) == nil { + return s + } + } + // 否则是 JSON 对象,重新序列化回字符串 + return string(raw) +} + func trimSlashes(s string) string { for len(s) > 0 && s[len(s)-1] == '/' { s = s[:len(s)-1] diff --git a/shortdrama/agent/react_agent.go b/shortdrama/agent/react_agent.go index 881c230..fc23f5d 100644 --- a/shortdrama/agent/react_agent.go +++ b/shortdrama/agent/react_agent.go @@ -81,6 +81,17 @@ func (a *ReActAgent) Run(ctx context.Context, userInput string) (string, error) continue } + if tc.Arguments == "" { + g.Log().Warningf(ctx, "ReAct step %d: 工具 %s 参数为空,跳过", step+1, tc.Name) + messages = append(messages, &ChatMessage{ + Role: RoleTool, + Content: "工具参数为空", + Name: tc.Name, + ToolCallID: tc.ID, + }) + continue + } + var args map[string]any if err := json.Unmarshal([]byte(tc.Arguments), &args); err != nil { g.Log().Warningf(ctx, "ReAct step %d: 参数解析失败: %v", step+1, err) diff --git a/shortdrama/consts/status.go b/shortdrama/consts/status.go index 63e64ae..5600f4c 100644 --- a/shortdrama/consts/status.go +++ b/shortdrama/consts/status.go @@ -7,6 +7,7 @@ const ( EpisodeStatusCompleted = "completed" EpisodeStatusFailed = "failed" + TaskStatusPending = "pending" TaskStatusGenerating = "generating" TaskStatusReview = "review" TaskStatusCompleted = "completed" diff --git a/shortdrama/controller/config_controller.go b/shortdrama/controller/config_controller.go index 984f0e2..8f8a0fa 100644 --- a/shortdrama/controller/config_controller.go +++ b/shortdrama/controller/config_controller.go @@ -2,7 +2,6 @@ package controller import ( "context" - "fmt" "video-factory/shortdrama/model/dto" "video-factory/shortdrama/model/entity" @@ -18,15 +17,6 @@ func (c *config) Get(ctx context.Context, req *dto.GetModelConfigReq) (res *dto. } func (c *config) Save(ctx context.Context, req *dto.SaveModelConfigReq) (res *struct{}, err error) { - if req.MinSingleDuration < 1 { - return nil, fmt.Errorf("单次生成最小时长必须大于等于1") - } - if req.MaxSingleDuration < 1 { - return nil, fmt.Errorf("单次生成最大时长必须大于等于1") - } - if req.MinSingleDuration > req.MaxSingleDuration { - return nil, fmt.Errorf("单次生成最小时长不能大于最大时长") - } return nil, service.ConfigService.Save(ctx, req) } diff --git a/shortdrama/dao/generation_task_dao.go b/shortdrama/dao/generation_task_dao.go index 044979f..a40aa60 100644 --- a/shortdrama/dao/generation_task_dao.go +++ b/shortdrama/dao/generation_task_dao.go @@ -23,6 +23,8 @@ func init() { episode_id INTEGER NOT NULL DEFAULT 0, -- 剧集ID segment_idx INTEGER NOT NULL DEFAULT 0, -- 段索引 status TEXT NOT NULL DEFAULT 'pending', -- 任务状态 + script TEXT NOT NULL DEFAULT '', -- 段脚本内容 + duration INTEGER NOT NULL DEFAULT 0, -- 段时长(秒) error_message TEXT NOT NULL DEFAULT '', created_at DATETIME DEFAULT (datetime('now','localtime')), updated_at DATETIME DEFAULT (datetime('now','localtime')) @@ -34,11 +36,6 @@ func init() { `ALTER TABLE `+public.TableNameGenerationTask+` ADD COLUMN segment_idx INTEGER NOT NULL DEFAULT 0`); err != nil { g.Log().Debugf(ctx, "添加 segment_idx 列失败(可能已存在): %v", err) } - // 迁移:添加 script_path 列 - if _, err := g.DB().Exec(ctx, - `ALTER TABLE `+public.TableNameGenerationTask+` ADD COLUMN script_path TEXT NOT NULL DEFAULT ''`); err != nil { - g.Log().Debugf(ctx, "添加 script_path 列失败(可能已存在): %v", err) - } // 迁移:添加 video_task_id 列 if _, err := g.DB().Exec(ctx, `ALTER TABLE `+public.TableNameGenerationTask+` ADD COLUMN video_task_id TEXT NOT NULL DEFAULT ''`); err != nil { @@ -69,6 +66,21 @@ func init() { `ALTER TABLE `+public.TableNameGenerationTask+` DROP COLUMN steps_data`); err != nil { g.Log().Debugf(ctx, "删除 steps_data 列失败(可能已不存在): %v", err) } + // 迁移:删除废弃的 script_path 列 + if _, err := g.DB().Exec(ctx, + `ALTER TABLE `+public.TableNameGenerationTask+` DROP COLUMN script_path`); err != nil { + g.Log().Debugf(ctx, "删除 script_path 列失败(可能已不存在): %v", err) + } + // 迁移:添加 script 列 + if _, err := g.DB().Exec(ctx, + `ALTER TABLE `+public.TableNameGenerationTask+` ADD COLUMN script TEXT NOT NULL DEFAULT ''`); err != nil { + g.Log().Debugf(ctx, "添加 script 列失败(可能已存在): %v", err) + } + // 迁移:添加 duration 列 + if _, err := g.DB().Exec(ctx, + `ALTER TABLE `+public.TableNameGenerationTask+` ADD COLUMN duration INTEGER NOT NULL DEFAULT 0`); err != nil { + g.Log().Debugf(ctx, "添加 duration 列失败(可能已存在): %v", err) + } } func (d *generationTaskDao) Insert(ctx context.Context, data *entity.GenerationTask) (id int64, err error) { diff --git a/shortdrama/dao/model_config_dao.go b/shortdrama/dao/model_config_dao.go index 38c002a..fc09254 100644 --- a/shortdrama/dao/model_config_dao.go +++ b/shortdrama/dao/model_config_dao.go @@ -27,8 +27,6 @@ func init() { video_api_key TEXT NOT NULL DEFAULT '', video_base_url TEXT NOT NULL DEFAULT '', video_model_name TEXT NOT NULL DEFAULT '', - max_single_duration INTEGER NOT NULL DEFAULT 15, - min_single_duration INTEGER NOT NULL DEFAULT 5, video_schema TEXT NOT NULL DEFAULT '', video_task_callback_url TEXT NOT NULL DEFAULT '', price_per_second INTEGER NOT NULL DEFAULT 0, @@ -40,6 +38,7 @@ func init() { // 清理旧字段 for _, col := range []string{ + "max_single_duration", "min_single_duration", "chat_provider", "video_provider", "video_no_duration_support", "video_model_category", "image_model_name", @@ -75,17 +74,15 @@ func init() { count, _ := g.DB().Model(public.TableNameModelConfig).Ctx(ctx).Count() if count == 0 { _, err := g.DB().Model(public.TableNameModelConfig).Ctx(ctx).Data(g.Map{ - "chat_api_key": "sk-ws-H.RXDDPMR.E8Rn.MEUCIQD00xjpXeNJlnXQejmbFTyC5ILHF8hLB2at4QXU6nf6GwIgae5TA8siaylUxMBDQcOWFfNx6hxB4rFsdCxoV6md810", - "chat_base_url": "https://dashscope.aliyuncs.com/compatible-mode", - "chat_model_name": "qwen-math-turbo", - "max_tokens": 8192, - "temperature": 0.7, - "video_api_key": "sk-ws-H.RXDDPMR.E8Rn.MEUCIQD00xjpXeNJlnXQejmbFTyC5ILHF8hLB2at4QXU6nf6GwIgae5TA8siaylUxMBDQcOWFfNx6hxB4rFsdCxoV6md810", - "video_base_url": "https://ws-0mn9581u5vt4lcx9.cn-beijing.maas.aliyuncs.com/api/v1/tasks", - "video_task_callback_url": "https://ws-0mn9581u5vt4lcx9.cn-beijing.maas.aliyuncs.com/api/v1/tasks/{task_id}", - "video_model_name": "wan2.7-r2v", - "max_single_duration": 5, - "min_single_duration": 2, + "chat_api_key": "", + "chat_base_url": "", + "chat_model_name": "", + "max_tokens": 4096, + "temperature": 0.85, + "video_api_key": "", + "video_base_url": "", + "video_task_callback_url": "", + "video_model_name": "", "price_per_second": 10, "created_at": gtime.Now().Format("Y-m-d H:i:s"), "updated_at": gtime.Now().Format("Y-m-d H:i:s"), diff --git a/shortdrama/model/dto/config_dto.go b/shortdrama/model/dto/config_dto.go index 9f8a16d..0aa3bc3 100644 --- a/shortdrama/model/dto/config_dto.go +++ b/shortdrama/model/dto/config_dto.go @@ -31,9 +31,7 @@ type SaveModelConfigReq struct { VideoApiKey string `v:"required" json:"videoApiKey" dc:"视频模型API密钥"` VideoBaseUrl string `v:"required|url" json:"videoBaseUrl" dc:"视频模型API接口地址"` VideoModelName string `v:"required" json:"videoModelName" dc:"视频模型名称"` - MaxSingleDuration int `v:"required" json:"maxSingleDuration" dc:"单段最大时长"` - MinSingleDuration int `v:"required" json:"minSingleDuration" dc:"单段最小时长"` - VideoSchema *gjson.Json `json:"videoSchema" dc:"视频生成模型schema"` + VideoSchema *gjson.Json `json:"videoSchema" dc:"视频生成模型schema(JSON, 含duration/ref/params/sizes)"` PricePerSecond int `v:"required|min:1" json:"pricePerSecond" dc:"每秒价格(分)"` VideoTaskCallbackUrl string `json:"videoTaskCallbackUrl" dc:"视频生成任务回调地址"` } diff --git a/shortdrama/model/entity/generation_task.go b/shortdrama/model/entity/generation_task.go index 5b3c9e9..0aa14fb 100644 --- a/shortdrama/model/entity/generation_task.go +++ b/shortdrama/model/entity/generation_task.go @@ -3,16 +3,17 @@ package entity import "github.com/gogf/gf/v2/os/gtime" type GenerationTask struct { - Id int64 `orm:"id" json:"id" dc:"任务ID" json:"id" json:"id" dc:"任务ID"` - DramaId int64 `orm:"drama_id" json:"dramaId" dc:"短剧ID" json:"drama_id" json:"dramaId" dc:"短剧ID"` - EpisodeId int64 `orm:"episode_id" json:"episodeId" dc:"剧集ID" json:"episode_id" json:"episodeId" dc:"剧集ID"` - SegmentIdx int `orm:"segment_idx" json:"segmentIdx" dc:"段索引" json:"segment_idx" json:"segmentIdx" dc:"段索引"` - Status string `orm:"status" json:"status" dc:"任务状态(pending/generating/review/completed/failed)" json:"status" json:"status" dc:"任务状态(pending/generating/review/completed/failed)"` - ScriptPath string `orm:"script_path" json:"scriptPath" dc:"脚本文件路径" json:"script_path" json:"scriptPath" dc:"脚本文件路径"` - VideoTaskId string `orm:"video_task_id" json:"videoTaskId" dc:"视频合成任务ID" json:"video_task_id" json:"videoTaskId" dc:"视频合成任务ID"` - VideoUrl string `orm:"video_url" json:"videoUrl" dc:"视频文件路径" json:"video_url" json:"videoUrl" dc:"视频文件路径"` - NumSegments int `orm:"num_segments" json:"numSegments" dc:"本集总段数" json:"num_segments" json:"numSegments" dc:"本集总段数"` - ErrorMessage string `orm:"error_message" json:"errorMessage" dc:"错误信息" json:"error_message" json:"errorMessage" dc:"错误信息"` - CreatedAt *gtime.Time `orm:"created_at" json:"createdAt" dc:"创建时间" json:"created_at" json:"createdAt" dc:"创建时间"` - UpdatedAt *gtime.Time `orm:"updated_at" json:"updatedAt" dc:"更新时间" json:"updated_at" json:"updatedAt" dc:"更新时间"` + Id int64 `orm:"id" json:"id" dc:"任务ID"` + DramaId int64 `orm:"drama_id" json:"dramaId" dc:"短剧ID"` + EpisodeId int64 `orm:"episode_id" json:"episodeId" dc:"剧集ID"` + SegmentIdx int `orm:"segment_idx" json:"segmentIdx" dc:"段索引"` + Status string `orm:"status" json:"status" dc:"任务状态(pending/generating/review/completed/failed)"` + Script string `orm:"script" json:"script" dc:"段脚本内容"` + Duration int `orm:"duration" json:"duration" dc:"段时长(秒)"` + VideoTaskId string `orm:"video_task_id" json:"videoTaskId" dc:"视频合成任务ID"` + VideoUrl string `orm:"video_url" json:"videoUrl" dc:"视频文件路径"` + NumSegments int `orm:"num_segments" json:"numSegments" dc:"本集总段数"` + ErrorMessage string `orm:"error_message" json:"errorMessage" dc:"错误信息"` + CreatedAt *gtime.Time `orm:"created_at" json:"createdAt" dc:"创建时间"` + UpdatedAt *gtime.Time `orm:"updated_at" json:"updatedAt" dc:"更新时间"` } diff --git a/shortdrama/model/entity/model_config.go b/shortdrama/model/entity/model_config.go index 4735c6d..c44954d 100644 --- a/shortdrama/model/entity/model_config.go +++ b/shortdrama/model/entity/model_config.go @@ -15,9 +15,7 @@ type ModelConfig struct { VideoApiKey string `orm:"video_api_key" json:"videoApiKey" dc:"视频模型API密钥"` VideoBaseUrl string `orm:"video_base_url" json:"videoBaseUrl" dc:"视频模型接口地址"` VideoModelName string `orm:"video_model_name" json:"videoModelName" dc:"视频模型名称"` - MaxSingleDuration int `orm:"max_single_duration" json:"maxSingleDuration" dc:"单段最大时长"` - MinSingleDuration int `orm:"min_single_duration" json:"minSingleDuration" dc:"单段最小时长"` - VideoSchema string `orm:"video_schema" json:"videoSchema" dc:"视频生成模型schema"` + VideoSchema string `orm:"video_schema" json:"videoSchema" dc:"视频生成模型schema(JSON, 含duration/ref/prams/sizes)"` PricePerSecond int `orm:"price_per_second" json:"pricePerSecond" dc:"每秒价格(分)"` VideoTaskCallbackUrl string `orm:"video_task_callback_url" json:"videoTaskCallbackUrl" dc:"视频生成任务回调地址"` CreatedAt *gtime.Time `orm:"created_at" json:"createdAt" dc:"创建时间"` diff --git a/shortdrama/model/segment_output.go b/shortdrama/model/segment_output.go index 438a760..13b482c 100644 --- a/shortdrama/model/segment_output.go +++ b/shortdrama/model/segment_output.go @@ -30,5 +30,5 @@ type SegmentCharacter struct { type VideoRef struct { Type string `json:"type"` // "character" | "scene" | "prop" Name string `json:"name"` // 实体名称 - MediaURL string `json:"mediaUrl"` // 参考素材 URL(base64 data URL 或 HTTP URL) + MediaURL string `json:"mediaUrl"` // 参考素材 URL(API 要求 HTTPS 公网 URL 或 OSS URL) } diff --git a/shortdrama/service/config_service.go b/shortdrama/service/config_service.go index ab16464..b1e4a90 100644 --- a/shortdrama/service/config_service.go +++ b/shortdrama/service/config_service.go @@ -2,10 +2,6 @@ package service import ( "context" - "encoding/json" - "io" - "net/http" - "strings" "video-factory/shortdrama/dao" "video-factory/shortdrama/model/dto" @@ -88,133 +84,14 @@ func (s *configService) SavePaymentConfig(ctx context.Context, cfg *entity.Payme return dao.PaymentConfigDao.Save(ctx, cfg) } -// syncModelDurationFromAPI 查询模型API,获取模型支持的参数和时长范围 +// syncModelDurationFromAPI 保存前校验配置完整性,设置默认值 func syncModelDurationFromAPI(ctx context.Context, cfg *entity.ModelConfig) { - if cfg.VideoApiKey == "" || cfg.VideoBaseUrl == "" || cfg.VideoModelName == "" { - return - } - - if cfg.MinSingleDuration <= 0 { - cfg.MinSingleDuration = 2 - g.Log().Infof(ctx, "使用默认单段最小时长: 2 秒") - } - if cfg.MaxSingleDuration <= 0 { - cfg.MaxSingleDuration = 15 - g.Log().Infof(ctx, "使用默认单段最大时长: 15 秒") - } if cfg.Temperature <= 0 { cfg.Temperature = 0.85 g.Log().Infof(ctx, "使用默认 Temperature: 0.85") } - - if cfg.ChatApiKey == "" || cfg.ChatBaseUrl == "" || cfg.ChatModelName == "" { - return - } - - modelsURL := strings.TrimRight(cfg.ChatBaseUrl, "/") + "/models" - chatReq, err := http.NewRequestWithContext(ctx, "GET", modelsURL, nil) - if err != nil { - g.Log().Warningf(ctx, "创建对话模型查询请求失败: %v", err) - return - } - chatReq.Header.Set("Authorization", "Bearer "+cfg.ChatApiKey) - chatReq.Header.Set("Content-Type", "application/json") - - chatResp, err := http.DefaultClient.Do(chatReq) - if err != nil { - g.Log().Warningf(ctx, "查询对话模型列表API失败: %v(不影响配置保存,max_tokens将使用默认值)", err) - return - } - defer chatResp.Body.Close() - - chatBody, _ := io.ReadAll(chatResp.Body) - var chatModelList struct { - Object string `json:"object"` - Data []struct { - Id string `json:"id"` - Object string `json:"object"` - Created int64 `json:"created"` - OwnedBy string `json:"owned_by"` - } `json:"data"` - FirstId string `json:"first_id"` - LastId string `json:"last_id"` - HasMore bool `json:"has_more"` - } - if err := json.Unmarshal(chatBody, &chatModelList); err != nil || chatModelList.Object != "list" { - g.Log().Debugf(ctx, "解析对话模型列表响应失败(可能非OpenAI兼容API): %v", err) - return - } - - chatModelName := strings.ToLower(cfg.ChatModelName) - if cfg.MaxTokens <= 0 { - cfg.MaxTokens = inferChatMaxTokens(chatModelName) - g.Log().Infof(ctx, "根据模型名称推断 max_tokens: %d(用户在表单中可手动修改)", cfg.MaxTokens) - } - if cfg.Temperature <= 0 { - cfg.Temperature = inferChatTemperature(chatModelName) - g.Log().Infof(ctx, "根据模型名称推断 Temperature: %.2f(用户在表单中可手动修改)", cfg.Temperature) - } -} - -func inferChatMaxTokens(modelName string) int { - switch { - case strings.Contains(modelName, "qwen-math"): - return 32768 - case strings.Contains(modelName, "qwen3"): - return 65536 - case strings.Contains(modelName, "qwen2.5"): - return 32768 - case strings.Contains(modelName, "qwen-max"): - return 8192 - case strings.Contains(modelName, "qwen-plus"): - return 16384 - case strings.Contains(modelName, "gpt-4o"): - return 16384 - case strings.Contains(modelName, "gpt-4"): - return 8192 - case strings.Contains(modelName, "gpt-3.5"): - return 4096 - case strings.Contains(modelName, "deepseek"): - return 8192 - case strings.Contains(modelName, "doubao") || strings.Contains(modelName, "豆包"): - return 65536 - case strings.Contains(modelName, "glm"): - return 8192 - case strings.Contains(modelName, "ernie"): - return 8192 - default: - return 4096 - } -} - -func inferChatTemperature(modelName string) float64 { - switch { - case strings.Contains(modelName, "qwen-math"): - return 0.5 - case strings.Contains(modelName, "qwen3"): - return 0.7 - case strings.Contains(modelName, "qwen2.5"): - return 0.7 - case strings.Contains(modelName, "qwen-max"): - return 0.8 - case strings.Contains(modelName, "qwen-plus"): - return 0.8 - case strings.Contains(modelName, "gpt-4o"): - return 0.8 - case strings.Contains(modelName, "gpt-4"): - return 0.7 - case strings.Contains(modelName, "gpt-3.5"): - return 0.7 - case strings.Contains(modelName, "deepseek"): - return 0.7 - case strings.Contains(modelName, "doubao") || strings.Contains(modelName, "豆包"): - return 0.95 - case strings.Contains(modelName, "glm"): - return 0.8 - case strings.Contains(modelName, "ernie"): - return 0.8 - default: - return 0.7 + cfg.MaxTokens = 4096 + g.Log().Infof(ctx, "使用默认 MaxTokens: 4096") } } diff --git a/shortdrama/service/drama_service.go b/shortdrama/service/drama_service.go index a93ccd4..1f66240 100644 --- a/shortdrama/service/drama_service.go +++ b/shortdrama/service/drama_service.go @@ -9,7 +9,10 @@ import ( "math" "net/http" "os" + "os/exec" "path/filepath" + "sort" + "strconv" "strings" "sync" "time" @@ -23,7 +26,6 @@ import ( "video-factory/shortdrama/model/dto" "video-factory/shortdrama/model/entity" - "github.com/Eyevinn/mp4ff/mp4" "github.com/gogf/gf/v2/database/gdb" "github.com/gogf/gf/v2/frame/g" "github.com/gogf/gf/v2/os/gcache" @@ -280,31 +282,16 @@ func (s *dramaService) GenerateEpisode(ctx context.Context, dramaId, epId int64, // 清理本集之前生成的 workspace 文件 cleanupEpisodeWorkspace(ctx, d.Title, ep.Index, ep.Title) - // 事务:删除旧任务 → 批量创建新任务 → 更新剧集状态 - taskIds := make([]int64, numSegments) + // 更新任务状态:pending → generating err = g.DB().Transaction(ctx, func(ctx context.Context, tx gdb.TX) error { - // 删除旧任务 - if _, e := tx.Model(public.TableNameGenerationTask).Ctx(ctx).Where("episode_id", epId).Delete(); e != nil { + // 将所有待生成任务标记为生成中 + if _, e := tx.Model(public.TableNameGenerationTask).Ctx(ctx). + Where("episode_id", epId). + Where("status", consts.TaskStatusPending). + Data(g.Map{"status": consts.TaskStatusGenerating}).Update(); e != nil { return e } - // 创建新任务 - for i := range segDurs { - segDur := segDurs[i] - r, e := tx.Model(public.TableNameGenerationTask).Ctx(ctx).Data(g.Map{ - "drama_id": dramaId, - "episode_id": epId, - "segment_idx": i, - "status": consts.TaskStatusGenerating, - "num_segments": numSegments, - }).Insert() - if e != nil { - return fmt.Errorf("创建第%d段任务失败: %w", i+1, e) - } - taskId, _ := r.LastInsertId() - taskIds[i] = taskId - g.Log().Infof(ctx, "第%d集第%d段任务创建(id=%d, segDur=%ds)", ep.Index, i+1, taskId, segDur) - } - // 更新剧集状态和生成模式 + // 更新剧集状态 _, e := tx.Model(public.TableNameEpisode).Ctx(ctx).Data(g.Map{"status": consts.EpisodeStatusGenerating, "video_url": ""}).Where("id", epId).Update() return e }) @@ -312,8 +299,14 @@ func (s *dramaService) GenerateEpisode(ctx context.Context, dramaId, epId int64, return err } - // 更新 poll cache + // 加载任务列表(状态已更新为 generating) tasks, _ := dao.GenerationTask.ListByEpisode(ctx, epId) + + // 构建 taskId 查找表 + taskIds := make(map[int]int64) + for _, t := range tasks { + taskIds[t.SegmentIdx] = t.Id + } setPollCache(ctx, epId, tasks) if mode == "serial" { @@ -351,7 +344,16 @@ func (s *dramaService) GenerateEpisode(ctx context.Context, dramaId, epId int64, wg.Add(1) sem <- struct{}{} go func(idx, dur int, tid int64) { - defer func() { <-sem; wg.Done() }() + defer func() { + if r := recover(); r != nil { + errMsg := fmt.Sprintf("panic: %v", r) + g.Log().Errorf(genCtx, "第%d集第%d段生成panic: %v", ep.Index, idx+1, r) + _ = dao.GenerationTask.UpdateFailed(genCtx, tid, errMsg) + clearPollCache(genCtx, epId) + } + <-sem + wg.Done() + }() startTime := time.Now() g.Log().Infof(genCtx, "第%d集第%d段开始生成(segDur=%ds)", ep.Index, idx+1, dur) if err := s.generateOneSegment(genCtx, d, ep, tid, idx, dur, "", genCtx2); err != nil { @@ -457,21 +459,83 @@ func (s *dramaService) generateOneSegment(ctx context.Context, d *entity.Drama, // 提交视频合成 // 构建引用列表:从Agent输出的角色名匹配预加载的参考图片 var videoRefs []model.VideoRef + + // 1. 角色参考 for _, ch := range segOutput.Characters { - if url := genCtx.LookupRef("演员", ch.Name); url != "" { + mediaURL := genCtx.LookupRef("演员", ch.Name) + if mediaURL == "" { + continue + } + videoRefs = append(videoRefs, model.VideoRef{ + Type: "character", Name: ch.Name, MediaURL: mediaURL, + }) + } + + // 2. 场景和道具参考:从脚本镜头中解析本段时间范围内涉及的场景和道具 + segEndTime := segStartTime + segDur + var coveredScenes []string + var coveredProps []string + if domain.IsShotsJSON(ep.Script) { + var shots []domain.Shot + if err := json.Unmarshal([]byte(ep.Script), &shots); err == nil { + for _, sh := range shots { + shStart := parseMMSSToSeconds(sh.StartTime) + shEnd := parseMMSSToSeconds(sh.EndTime) + // 检查镜头是否与本段时间范围重叠 + if shEnd <= segStartTime || shStart >= segEndTime { + continue + } + if sh.Scene != "" { + coveredScenes = append(coveredScenes, sh.Scene) + } + for _, p := range sh.Props { + if p != "" { + coveredProps = append(coveredProps, p) + } + } + } + } + } + + // 去重后添加场景参考 + seenScene := make(map[string]bool) + for _, name := range coveredScenes { + if seenScene[name] { + continue + } + seenScene[name] = true + if url := genCtx.LookupRef("场景", name); url != "" { videoRefs = append(videoRefs, model.VideoRef{ - Type: "character", Name: ch.Name, MediaURL: url, + Type: "scene", Name: name, MediaURL: url, }) } } - taskID, submitErr := s.submitVideoTask(ctx, d, ep, segIdx, segDur, segOutput.Scenes, videoRefs) + // 去重后添加道具参考 + seenProp := make(map[string]bool) + for _, name := range coveredProps { + if seenProp[name] { + continue + } + seenProp[name] = true + if url := genCtx.LookupRef("道具", name); url != "" { + videoRefs = append(videoRefs, model.VideoRef{ + Type: "prop", Name: name, MediaURL: url, + }) + } + } + + g.Log().Infof(ctx, "第%d段视频参考素材: %d个(角色%d 场景%d 道具%d)", + segIdx+1, len(videoRefs), len(segOutput.Characters), len(coveredScenes), len(coveredProps)) + + taskID, _, submitErr := s.submitVideoTask(ctx, d, ep, segIdx, segDur, segOutput.Scenes, videoRefs) if submitErr != nil { - g.Log().Warningf(ctx, "第%d段视频提交失败,后台轮询器将自动重试: %v", segIdx+1, submitErr) + g.Log().Errorf(ctx, "第%d段视频提交失败: %v", segIdx+1, submitErr) + return submitErr } else { g.Log().Infof(ctx, "第%d集第%d段视频已提交(taskId=%s, duration=%ds)", ep.Index, segIdx+1, taskID, segDur) } - // 更新 task 字段 + // 只更新 video_task_id 和 num_segments,不覆盖 script(保留 createPendingTasks 的原始 bodyJSON) updateFields := g.Map{ "num_segments": numSegments, "updated_at": nil, @@ -755,7 +819,7 @@ func (s *dramaService) ContinueSegment(ctx context.Context, taskId int64) error return nil } -// FeedbackSegment 反馈并重新生成本段 +// FeedbackSegment 反馈并重新提交本段视频(跳过 Agent,直接使用 task.Script 重新调用模型 API) func (s *dramaService) FeedbackSegment(ctx context.Context, taskId int64, feedback string) error { task, err := dao.GenerationTask.GetOne(ctx, taskId) if err != nil { @@ -780,6 +844,11 @@ func (s *dramaService) FeedbackSegment(ctx context.Context, taskId int64, feedba return err } + modelCfg := ConfigService.Get(ctx) + if modelCfg.VideoApiKey == "" || modelCfg.VideoBaseUrl == "" { + return fmt.Errorf("视频模型未配置") + } + d, err := dao.Drama.GetOne(ctx, task.DramaId) if err != nil || d == nil { return fmt.Errorf("短剧不存在") @@ -789,29 +858,49 @@ func (s *dramaService) FeedbackSegment(ctx context.Context, taskId int64, feedba return fmt.Errorf("剧集不存在") } - modelCfg := ConfigService.Get(ctx) - - segDurs := calcSegDurs(d.EpisodeDuration, modelCfg) - segDur := segDurs[task.SegmentIdx] - - // 预加载引用数据 - feedbackGenCtx, fbErr := BuildGenerationContext(context.Background(), d) - if fbErr != nil { - g.Log().Errorf(ctx, "反馈重试构建生成上下文失败: %v", fbErr) - return fbErr - } - go func() { genCtx := context.Background() genCtx = agent.WithDramaID(genCtx, d.Id) - g.Log().Infof(genCtx, "第%d集第%d段根据反馈重新生成", ep.Index, task.SegmentIdx+1) - if err := s.generateOneSegment(genCtx, d, ep, taskId, task.SegmentIdx, segDur, feedback, feedbackGenCtx); err != nil { - g.Log().Errorf(genCtx, "第%d集第%d段重新生成失败: %v", ep.Index, task.SegmentIdx+1, err) - _ = dao.GenerationTask.UpdateFailed(genCtx, taskId, err.Error()) + g.Log().Infof(genCtx, "第%d集第%d段重新提交视频(跳过Agent)", ep.Index, task.SegmentIdx+1) + + // 将 feedback 追加到 task.Script 的 prompt 末尾,不覆盖原有内容 + // 同时用当前数据库中的模型名更新 model 字段 + bodyBytes := []byte(task.Script) + { + var bodyMap map[string]any + if err := json.Unmarshal(bodyBytes, &bodyMap); err == nil { + if input, ok := bodyMap["input"].(map[string]any); ok { + if p, ok := input["prompt"].(string); ok && feedback != "" { + input["prompt"] = p + "\n\n【修改意见】" + feedback + } + } + // 用当前配置的模型名覆盖(兼容数据库配置变更后重做) + bodyMap["model"] = modelCfg.VideoModelName + bodyBytes, _ = json.Marshal(bodyMap) + } + } + + newTaskID, resolvedBody, submitErr := resubmitVideoTask(genCtx, modelCfg.VideoApiKey, modelCfg.VideoBaseUrl, bodyBytes) + if submitErr != nil { + g.Log().Errorf(genCtx, "第%d集第%d段重新提交视频失败: %v", ep.Index, task.SegmentIdx+1, submitErr) + _ = dao.GenerationTask.UpdateFailed(genCtx, taskId, submitErr.Error()) clearPollCache(genCtx, task.EpisodeId) if ts, e := dao.GenerationTask.ListByEpisode(genCtx, task.EpisodeId); e == nil { setPollCache(genCtx, task.EpisodeId, ts) } + return + } + + g.Log().Infof(genCtx, "第%d集第%d段重提提交成功(taskId=%s)", ep.Index, task.SegmentIdx+1, newTaskID) + _ = dao.GenerationTask.UpdateFields(genCtx, taskId, g.Map{ + "video_task_id": newTaskID, + "script": string(resolvedBody), + "updated_at": nil, + }) + clearPollCache(genCtx, task.EpisodeId) + _ = dao.Episode.UpdateStatus(genCtx, task.EpisodeId, consts.EpisodeStatusGenerating, "") + if ts, e := dao.GenerationTask.ListByEpisode(genCtx, task.EpisodeId); e == nil { + setPollCache(genCtx, task.EpisodeId, ts) } }() @@ -989,6 +1078,66 @@ func (s *dramaService) pollPendingVideos(ctx context.Context) { continue } + // 任务没有 video_task_id —— 从任务自带 script 或剧集脚本重试提交 + ep, _ := dao.Episode.GetOne(ctx, task.EpisodeId) + d, _ := dao.Drama.GetOne(ctx, task.DramaId) + if ep != nil && d != nil && (task.Script != "" || ep.Script != "") { + // 优先使用任务自带的 script(新格式为 JSON 请求体),旧数据回退到 ep.Script + scriptText := task.Script + if scriptText != "" { + // 新格式:task.Script 为 JSON,从中提取 prompt + var bodyMap map[string]any + if err := json.Unmarshal([]byte(scriptText), &bodyMap); err == nil { + // 如果是完整的 bodyJSON(含 model/input/parameters),说明是 createPendingTasks 创建的新任务 + // 尚未经过 generateOneSegment 的 Agent 处理,跳过重试以避免重复提交 + if _, ok := bodyMap["model"]; ok { + if _, ok := bodyMap["parameters"]; ok { + g.Log().Infof(ctx, "轮询器: 任务 %d 第%d段为 bodyJSON 格式,跳过重试等待主流程处理", + task.Id, task.SegmentIdx+1) + continue + } + } + // 尝试从 bodyJSON 的 input.prompt 提取(兼容旧的重试场景) + if p, ok := bodyMap["prompt"].(string); ok && p != "" { + scriptText = p + } else if input, ok := bodyMap["input"].(map[string]any); ok { + if p, ok := input["prompt"].(string); ok && p != "" { + scriptText = p + } + } + } + // 如果不是 JSON 或没有 prompt 字段,保持原样(旧格式纯文本) + } + if scriptText == "" { + scriptText = ep.Script + } + scenes := []model.SegmentScene{{Description: scriptText}} + g.Log().Infof(ctx, "轮询器: 任务 %d 第%d段尝试重提交视频", + task.Id, task.SegmentIdx+1) + // 优先使用任务自带的 duration,旧数据回退到 calcSegDurs + segD := task.Duration + if segD <= 0 { + segDur := calcSegDurs(d.EpisodeDuration, modelCfg) + segD = 10 + if len(segDur) > task.SegmentIdx { + segD = segDur[task.SegmentIdx] + } + } + taskID, _, err := s.submitVideoTask(ctx, d, ep, task.SegmentIdx, segD, scenes, nil) + if err != nil { + g.Log().Errorf(ctx, "轮询器: 任务 %d 第%d段视频重提交失败: %v", + task.Id, task.SegmentIdx+1, err) + } else { + g.Log().Infof(ctx, "轮询器: 任务 %d 第%d段视频重提交成功(taskId=%s)", + task.Id, task.SegmentIdx+1, taskID) + _ = dao.GenerationTask.UpdateFields(ctx, task.Id, g.Map{"video_task_id": taskID}) + epUpdates[task.EpisodeId] = true + } + } else { + g.Log().Warningf(ctx, "轮询器: 任务 %d 第%d段无剧集脚本,%s", + task.Id, task.SegmentIdx+1, "等待主流程完成") + } + } // 刷新所有 generating 任务的 episode 缓存(不管状态有无变化,确保前端轮询不走 DB) @@ -1005,9 +1154,21 @@ func (s *dramaService) pollPendingVideos(ctx context.Context) { } +// buildPollURL 从视频创建 URL 构建任务查询 URL(DashScope API 规范) +// 创建端点: /api/v1/services/aigc/video-generation/video-synthesis +// 查询端点: /api/v1/tasks/{taskId} +func buildPollURL(baseURL, taskId string) string { + // 查找 "/api/v1/" 路径并替换后面的内容为 tasks/{taskId} + if idx := strings.Index(baseURL, "/api/v1/"); idx > 0 { + return baseURL[:idx] + "/api/v1/tasks/" + taskId + } + // fallback:直接拼接 + return strings.TrimRight(baseURL, "/") + "/" + taskId +} + // pollVideoTaskOnce 单次查询视频生成任务状态 func (s *dramaService) pollVideoTaskOnce(ctx context.Context, modelCfg *entity.ModelConfig, taskId string) (string, error) { - queryURL := strings.TrimRight(modelCfg.VideoBaseUrl, "/") + "/" + taskId + queryURL := buildPollURL(modelCfg.VideoBaseUrl, taskId) req, err := http.NewRequestWithContext(ctx, "GET", queryURL, nil) if err != nil { @@ -1062,17 +1223,17 @@ func (s *dramaService) pollVideoTaskOnce(ctx context.Context, modelCfg *entity.M if videoURL != "" { return videoURL, nil } + // SUCCEEDED 但没有视频URL,输出完整响应帮助调试 + g.Log().Warningf(ctx, "任务 %s 状态为 SUCCEEDED 但未返回视频URL,完整响应: %s", taskId, string(data)) + return "", fmt.Errorf("SUCCEEDED但视频URL为空,完整响应: %s", string(data)) case videoTaskFailed: return "", fmt.Errorf("FAILED: %s", buildErrMsg()) case videoTaskRunning, videoTaskPending: return "", fmt.Errorf("RUNNING") } - errMsg := buildErrMsg() - if errMsg != result.Output.TaskStatus { - return "", fmt.Errorf("任务状态: %s", errMsg) - } - return "", fmt.Errorf("视频URL为空") + g.Log().Warningf(ctx, "任务 %s 返回未知状态,完整响应: %s", taskId, string(data)) + return "", fmt.Errorf("未知任务状态: %s,完整响应: %s", result.Output.TaskStatus, string(data)) } // mapVideoTaskStatus 将供应商任务状态映射为内部状态 @@ -1093,11 +1254,12 @@ func mapVideoTaskStatus(s string) videoTaskStatus { // ==================== Video Generation ==================== -// submitVideoTask 提交视频合成任务 -func (s *dramaService) submitVideoTask(ctx context.Context, d *entity.Drama, ep *entity.Episode, segIdx, segDur int, scenes []model.SegmentScene, refs []model.VideoRef) (string, error) { +// submitVideoTask 提交视频合成任务,返回 (taskID, requestBodyJSON, error) +// requestBodyJSON 是调用视频模型 API 时发送的完整请求体 JSON 字符串 +func (s *dramaService) submitVideoTask(ctx context.Context, d *entity.Drama, ep *entity.Episode, segIdx, segDur int, scenes []model.SegmentScene, refs []model.VideoRef) (string, string, error) { modelCfg := ConfigService.Get(ctx) if modelCfg.VideoApiKey == "" || modelCfg.VideoBaseUrl == "" || modelCfg.VideoModelName == "" { - return "", fmt.Errorf("视频模型未配置") + return "", "", fmt.Errorf("视频模型未配置") } if len(scenes) == 0 { @@ -1110,111 +1272,222 @@ func (s *dramaService) submitVideoTask(ctx context.Context, d *entity.Drama, ep } } + // 从 video_schema 读取模型特定配置(结构按 API 请求格式: input / parameters) + refMaxCount := 5 // 默认最多5个 + promptMaxChars := 0 + if modelCfg.VideoSchema != "" { + var vs map[string]any + if err := json.Unmarshal([]byte(modelCfg.VideoSchema), &vs); err == nil { + refMaxCount = intVal(nested(vs, "input", "reference_urls", "total_max"), refMaxCount) + promptMaxChars = intVal(nested(vs, "input", "prompt", "max_chars"), 0) + } + } + sceneDescs := make([]string, 0, len(scenes)) for _, s := range scenes { sceneDescs = append(sceneDescs, s.Description) } - sceneText := strings.Join(sceneDescs, ";") - // 构建参考素材说明 - var refImages []string - refPrompt := "" - if len(refs) > 0 { - refImages = make([]string, 0, len(refs)) - var refParts []string - for _, r := range refs { - refParts = append(refParts, fmt.Sprintf("%s(%s)", r.Name, r.Type)) - if r.MediaURL != "" { - refImages = append(refImages, r.MediaURL) - } + // 构建 reference_urls 数组:按 演员→场景→道具 优先级排序 + type namedURL struct { + url string + name string + } + var namedURLs []namedURL + + var chars, sceneRefs, props []model.VideoRef + for _, r := range refs { + switch r.Type { + case "character": + chars = append(chars, r) + case "scene": + sceneRefs = append(sceneRefs, r) + case "prop": + props = append(props, r) + } + } + orderedRefs := append(append(chars, sceneRefs...), props...) + + for _, r := range orderedRefs { + if len(namedURLs) >= refMaxCount { + break + } + if r.MediaURL != "" { + namedURLs = append(namedURLs, namedURL{url: r.MediaURL, name: r.Name}) } - refPrompt = fmt.Sprintf(",参考素材:%s", strings.Join(refParts, "、")) } - prompt := fmt.Sprintf("短剧《%s》第%d集第%d段:%s%s", d.Title, ep.Index, segIdx+1, sceneText, refPrompt) - maxInputLen := 3000 - if modelCfg.MaxTokens > 0 { - maxInputLen = modelCfg.MaxTokens + // 构建 prompt:将场景描述中的实体名称替换为 character1/character2/... 直接引用参考素材 + refURLs := make([]string, len(namedURLs)) + refExplanation := "" + if len(namedURLs) > 0 { + type nameLabel struct { + name string + label string + } + nlList := make([]nameLabel, len(namedURLs)) + for i, nu := range namedURLs { + label := fmt.Sprintf("character%d", i+1) + nlList[i] = nameLabel{name: nu.name, label: label} + refURLs[i] = nu.url + } + sort.Slice(nlList, func(i, j int) bool { + return len(nlList[i].name) > len(nlList[j].name) + }) + // 替换场景描述中的实体名为 characterN(模型仅通过此方式识别参考角色) + for _, nl := range nlList { + sceneText = strings.ReplaceAll(sceneText, nl.name, nl.label) + } + // 构建角色引用说明(按 character1/2/3 顺序) + var refParts []string + for i, nu := range namedURLs { + refParts = append(refParts, fmt.Sprintf("character%d=%s", i+1, nu.name)) + } + refExplanation = fmt.Sprintf("。角色引用说明:%s", strings.Join(refParts, "、")) } - if len([]rune(prompt)) > maxInputLen { + + prompt := fmt.Sprintf("短剧《%s》第%d集第%d段:%s%s", d.Title, ep.Index, segIdx+1, sceneText, refExplanation) + if promptMaxChars > 0 && len([]rune(prompt)) > promptMaxChars { prefix := fmt.Sprintf("短剧《%s》第%d集第%d段:", d.Title, ep.Index, segIdx+1) - keepLen := maxInputLen - len([]rune(prefix)) - if keepLen < 0 { - keepLen = 0 + refLen := len([]rune(refExplanation)) + keepSceneLen := promptMaxChars - refLen - len([]rune(prefix)) + if keepSceneLen < 50 { + keepSceneLen = 50 } sceneRunes := []rune(sceneText) - if keepLen < len(sceneRunes) { - sceneText = "..." + string(sceneRunes[len(sceneRunes)-keepLen+3:]) + if keepSceneLen < len(sceneRunes) { + sceneText = "..." + string(sceneRunes[len(sceneRunes)-keepSceneLen+3:]) } - prompt = prefix + sceneText - g.Log().Infof(ctx, "prompt超长已截断至%d字符(原%d字符)", maxInputLen, len([]rune(prompt))) + prompt = prefix + sceneText + refExplanation + g.Log().Infof(ctx, "prompt超长已截断至%d字符(原%d字符)", promptMaxChars, len([]rune(prompt))) } - taskId, err := createVideoTask(ctx, modelCfg.VideoApiKey, modelCfg.VideoBaseUrl, modelCfg.VideoModelName, - prompt, refImages, segDur, resolveVideoSize(d.Resolution, d.AspectRatio)) + // 读取 negative_prompt.md 文件作为反向提示词 + negativePrompt := "" + if data, err := os.ReadFile("negative_prompt.md"); err == nil { + if np := strings.TrimSpace(string(data)); np != "" { + negativePrompt = np + } + } + + taskId, requestJSON, err := createVideoTask(ctx, modelCfg.VideoApiKey, modelCfg.VideoBaseUrl, modelCfg.VideoModelName, + prompt, negativePrompt, refURLs, segDur, d.Resolution, d.AspectRatio, modelCfg.VideoSchema) if err != nil { errStr := strings.ToLower(err.Error()) if strings.Contains(errStr, "duration") && (strings.Contains(errStr, "not support") || strings.Contains(errStr, "not supported")) { - taskId, err = createVideoTask(ctx, modelCfg.VideoApiKey, modelCfg.VideoBaseUrl, modelCfg.VideoModelName, - prompt, refImages, 0, resolveVideoSize(d.Resolution, d.AspectRatio)) + taskId, requestJSON, err = createVideoTask(ctx, modelCfg.VideoApiKey, modelCfg.VideoBaseUrl, modelCfg.VideoModelName, + prompt, negativePrompt, refURLs, 0, d.Resolution, d.AspectRatio, modelCfg.VideoSchema) } if err != nil { - return "", fmt.Errorf("视频合成请求失败: %w", err) + return "", "", fmt.Errorf("视频合成请求失败: %w", err) } } if taskId == "" { - return "", fmt.Errorf("视频合成任务ID为空") + return "", "", fmt.Errorf("视频合成任务ID为空") } - return taskId, nil + return taskId, requestJSON, nil } -// createVideoTask 调用视频生成API提交任务 -func createVideoTask(ctx context.Context, apiKey, baseURL, modelName, prompt string, images []string, duration int, size string) (string, error) { - var ( - inputImages []string - paramsSize string - paramsDur int - hasParams bool - ) - if len(images) > 0 { - inputImages = images - } - if size != "" { - paramsSize = size - hasParams = true - } - if duration > 0 { - paramsDur = duration - hasParams = true - } - +// createVideoTask 调用视频生成API提交任务,返回 (taskID, requestBodyJSON, error) +// API 请求体结构由 model_config.video_schema 定义,包含: +// +// params: 额外参数(如 audio/shot_type/watermark 等) +// sizes: 分辨率→尺寸字符串映射表 +func createVideoTask(ctx context.Context, apiKey, baseURL, modelName, prompt, negativePrompt string, refURLs []string, duration int, resolution, aspectRatio, videoSchema string) (string, string, error) { body := map[string]any{ "model": modelName, "input": map[string]any{ "prompt": prompt, }, } - if len(inputImages) > 0 { - body["input"].(map[string]any)["images"] = inputImages + + // input.negative_prompt — 反向提示词(从 negative_prompt.md 读取,不为空时才发送) + if negativePrompt != "" { + body["input"].(map[string]any)["negative_prompt"] = negativePrompt } - if hasParams { - p := map[string]any{} - if paramsSize != "" { - p["size"] = paramsSize + + // input.reference_urls — 纯 URL 字符串数组,顺序对应 character1/character2/... + // 解析 refURLs:http/https/data: 开头的保持原样,文件路径转为 base64 data URL + for i, u := range refURLs { + if strings.HasPrefix(u, "http://") || strings.HasPrefix(u, "https://") || strings.HasPrefix(u, "data:") { + continue } - if paramsDur > 0 { - p["duration"] = paramsDur + if b64, err := imageFileToBase64(u); err == nil { + refURLs[i] = b64 + } else { + g.Log().Warningf(ctx, "reference_urls 元素无法解析为文件或URL,跳过: %s", u) } - body["parameters"] = p } + if len(refURLs) > 0 { + body["input"].(map[string]any)["reference_urls"] = refURLs + } else { + // 无参考图时使用默认首帧 + if imgData, err := os.ReadFile(DefaultFirstFramePath); err == nil { + b64 := base64.StdEncoding.EncodeToString(imgData) + body["input"].(map[string]any)["reference_urls"] = []string{"data:image/png;base64," + b64} + g.Log().Infof(ctx, "使用默认首帧图作为 reference_urls") + } else { + g.Log().Warningf(ctx, "读取默认首帧图失败: %v", err) + } + } + + // 从 video_schema 读取模型特定参数(结构按 API 请求格式: input / parameters) + params := map[string]any{} + if videoSchema != "" { + var vs map[string]any + if err := json.Unmarshal([]byte(videoSchema), &vs); err == nil { + // 尺寸映射 + if sizes := nested(vs, "parameters", "size", "sizes"); sizes != nil { + if sm, ok := sizes.(map[string]any); ok { + if v := resolveSizeFromSchema(sm, resolution, aspectRatio); v != "" { + params["size"] = v + } + } + } + // 遍历 parameters 读取各字段默认值注入 API 请求,跳过 supported_models 限制的参数 + if paramsDef := nested(vs, "parameters"); paramsDef != nil { + if pm, ok := paramsDef.(map[string]any); ok { + for k, v := range pm { + paramDef, ok := v.(map[string]any) + if !ok { + continue + } + // 如果有 supported_models 且不包含当前模型,跳过 + if sm, ok := paramDef["supported_models"]; ok { + if models, ok := sm.([]any); ok { + found := false + for _, m := range models { + if ms, ok := m.(string); ok && ms == modelName { + found = true + break + } + } + if !found { + continue + } + } + } + if def, ok := paramDef["default"]; ok { + params[k] = def + } + } + } + } + } + } + + if duration > 0 { + params["duration"] = duration + } + body["parameters"] = params payload, _ := json.Marshal(body) req, err := http.NewRequestWithContext(ctx, "POST", baseURL, strings.NewReader(string(payload))) if err != nil { - return "", fmt.Errorf("创建请求失败: %w", err) + return "", "", fmt.Errorf("创建请求失败: %w", err) } req.Header.Set("Authorization", "Bearer "+apiKey) req.Header.Set("Content-Type", "application/json") @@ -1223,7 +1496,7 @@ func createVideoTask(ctx context.Context, apiKey, baseURL, modelName, prompt str client := &http.Client{Timeout: 30 * time.Second} resp, err := client.Do(req) if err != nil { - return "", fmt.Errorf("请求失败: %w", err) + return "", "", fmt.Errorf("请求失败: %w", err) } defer resp.Body.Close() @@ -1236,40 +1509,125 @@ func createVideoTask(ctx context.Context, apiKey, baseURL, modelName, prompt str Message string `json:"message"` } if err := json.Unmarshal(respData, &result); err != nil { - return "", fmt.Errorf("解析响应失败: %s", string(respData)) + return "", "", fmt.Errorf("解析响应失败: %s", string(respData)) } if result.Code != "" { - return "", fmt.Errorf("请求失败(code=%s): %s", result.Code, string(respData)) + return "", "", fmt.Errorf("请求失败(code=%s): %s", result.Code, string(respData)) } if result.Output.TaskID == "" { - return "", fmt.Errorf("任务ID为空") + return "", "", fmt.Errorf("任务ID为空") } - return result.Output.TaskID, nil + return result.Output.TaskID, string(payload), nil } -// resolveVideoSize 根据分辨率和宽高比计算视频尺寸 -func resolveVideoSize(resolution, aspectRatio string) string { - res := resolution - if res == "" { - res = "720P" +// resubmitVideoTask 使用已有的 bodyJSON 重新提交视频任务(跳过 Agent,直接调 API) +func resubmitVideoTask(ctx context.Context, apiKey, baseURL string, bodyJSON []byte) (string, []byte, error) { + var bodyMap map[string]any + if err := json.Unmarshal(bodyJSON, &bodyMap); err != nil { + return "", nil, fmt.Errorf("解析 bodyJSON 失败: %w", err) } - if strings.Contains(aspectRatio, "16:9") { - if res == "1080P" { - return "1920*1080" + + // 解析 reference_urls:http/https/data: 开头的保持原样,文件路径转为 base64 + if input, ok := bodyMap["input"].(map[string]any); ok { + if refs, ok := input["reference_urls"].([]any); ok { + resolved := make([]any, len(refs)) + for i, r := range refs { + u, _ := r.(string) + if strings.HasPrefix(u, "http://") || strings.HasPrefix(u, "https://") || strings.HasPrefix(u, "data:") { + resolved[i] = u + } else if b64, err := imageFileToBase64(u); err == nil { + resolved[i] = b64 + } else { + g.Log().Warningf(ctx, "reference_urls 元素无法解析: %s", u) + resolved[i] = u + } + } + input["reference_urls"] = resolved } - return "1280*720" } - if strings.Contains(aspectRatio, "1:1") { - if res == "1080P" { - return "1080*1080" + + payload, _ := json.Marshal(bodyMap) + req, err := http.NewRequestWithContext(ctx, "POST", baseURL, strings.NewReader(string(payload))) + if err != nil { + return "", nil, fmt.Errorf("创建请求失败: %w", err) + } + req.Header.Set("Authorization", "Bearer "+apiKey) + req.Header.Set("Content-Type", "application/json") + req.Header.Set("X-DashScope-Async", "enable") + + client := &http.Client{Timeout: 30 * time.Second} + resp, err := client.Do(req) + if err != nil { + return "", nil, fmt.Errorf("请求失败: %w", err) + } + defer resp.Body.Close() + + respData, _ := io.ReadAll(resp.Body) + var result struct { + Output struct { + TaskID string `json:"task_id"` + } `json:"output"` + Code string `json:"code"` + Message string `json:"message"` + } + if err := json.Unmarshal(respData, &result); err != nil { + return "", nil, fmt.Errorf("解析响应失败: %s", string(respData)) + } + if result.Code != "" { + return "", nil, fmt.Errorf("请求失败(code=%s): %s", result.Code, string(respData)) + } + if result.Output.TaskID == "" { + return "", nil, fmt.Errorf("任务ID为空") + } + return result.Output.TaskID, payload, nil +} + +// resolveSizeFromSchema 从 video_schema.sizes 中根据分辨率和宽高比查找尺寸字符串 +// sizes 结构示例: {"720P":{"9:16":"720*1280","16:9":"1280*720"}} +func resolveSizeFromSchema(sizes map[string]any, resolution, aspectRatio string) string { + if resolution == "" || aspectRatio == "" { + return "" + } + res := strings.ToUpper(resolution) + ar := strings.ReplaceAll(aspectRatio, ":", ":") + resMap, ok := sizes[res].(map[string]any) + if !ok { + return "" + } + if v, ok := resMap[ar].(string); ok && v != "" { + return v + } + return "" +} + +// nested 从嵌套 map 中按 keys 路径安全取值 +func nested(m map[string]any, keys ...string) any { + for i, k := range keys { + if m == nil { + return nil + } + if i == len(keys)-1 { + return m[k] + } + if v, ok := m[k].(map[string]any); ok { + m = v + } else { + return nil } - return "720*720" } - // 默认竖屏 9:16 - if res == "1080P" { - return "1080*1920" + return nil +} + +// intVal 将 any 类型转为 int,非数值返回 fallback +func intVal(v any, fallback int) int { + switch n := v.(type) { + case float64: + return int(n) + case int: + return n + default: + return fallback } - return "720*1280" } func (s *dramaService) mergeEpisodeVideo(ctx context.Context, epId int64, lastTaskId int64) error { @@ -1351,6 +1709,14 @@ func (s *dramaService) mergeEpisodeVideo(ctx context.Context, epId int64, lastTa } } + // 叠加背景音乐(如果存在) + if bgmPath, err := s.overlayBackgroundMusic(ctx, d.Id, finalPath); err == nil { + finalPath = bgmPath + g.Log().Infof(ctx, "已叠加背景音乐: %s", finalPath) + } else { + g.Log().Warningf(ctx, "叠加背景音乐失败(跳过): %v", err) + } + // 事务:更新最后一段 task 状态 + 剧集状态和 video_url err = g.DB().Transaction(ctx, func(ctx context.Context, tx gdb.TX) error { // 更新最后一段任务为 completed @@ -1385,6 +1751,55 @@ func (s *dramaService) mergeEpisodeVideo(ctx context.Context, epId int64, lastTa // ==================== Video Helpers ==================== +// overlayBackgroundMusic 使用 ffmpeg 为视频叠加背景音乐(循环混音) +func (s *dramaService) overlayBackgroundMusic(ctx context.Context, dramaId int64, videoPath string) (string, error) { + bgmList, err := dao.BackgroundMusic.ListByDrama(ctx, dramaId) + if err != nil || len(bgmList) == 0 { + return "", fmt.Errorf("无背景音乐配置") + } + bgmPath := bgmList[0].FilePath + if bgmPath == "" { + return "", fmt.Errorf("背景音乐文件路径为空") + } + if _, err := os.Stat(bgmPath); os.IsNotExist(err) { + return "", fmt.Errorf("背景音乐文件不存在: %s", bgmPath) + } + // 检查 ffmpeg 是否可用 + if _, err := exec.LookPath("ffmpeg"); err != nil { + return "", fmt.Errorf("ffmpeg 不可用: %w", err) + } + + ext := filepath.Ext(videoPath) + bgmOutput := strings.TrimSuffix(videoPath, ext) + "_bgm" + ext + // ffmpeg 命令:循环 BGM 并降低音量,与视频原音频混音 + cmd := exec.Command("ffmpeg", + "-stream_loop", "-1", + "-i", bgmPath, + "-i", videoPath, + "-filter_complex", "[0:a]volume=0.15[bgm];[1:a][bgm]amix=inputs=2:duration=first[audio]", + "-map", "1:v", + "-map", "[audio]", + "-c:v", "copy", + "-shortest", + "-y", + bgmOutput, + ) + if output, err := cmd.CombinedOutput(); err != nil { + return "", fmt.Errorf("ffmpeg 失败: %w, 输出: %s", err, string(output)) + } + + // 替换原文件 + if err := os.Rename(bgmOutput, videoPath); err != nil { + // 如果替换失败,尝试移除原文件后重命名 + os.Remove(videoPath) + if err2 := os.Rename(bgmOutput, videoPath); err2 != nil { + os.Remove(bgmOutput) + return "", fmt.Errorf("替换视频文件失败: %w", err2) + } + } + return videoPath, nil +} + // waitForSegmentVideo 等待指定任务的视频生成完成。 // 轮询视频 API 直到视频就绪,下载到本地,更新 DB 的 video_url。 // 用于串行模式中让下一段能提取上一段的尾帧作为首帧。 @@ -1438,231 +1853,36 @@ func (s *dramaService) concatVideos(inputs []string, output string) error { return fmt.Errorf("需要至少2个输入文件才能合并") } - // 1. 读取所有 MP4 文件到内存 - var files []*mp4.File + // 创建 ffmpeg concat demuxer 文件列表 + filelist := output + ".filelist.txt" + var lines []string for _, path := range inputs { - f, err := mp4.ReadMP4File(path) + absPath, err := filepath.Abs(path) if err != nil { - return fmt.Errorf("读取 %s 失败: %w", path, err) + absPath = path } - files = append(files, f) + // 转义单引号 + escaped := strings.ReplaceAll(absPath, "'", "'\\''") + lines = append(lines, "file '"+escaped+"'") } - - // 2. 计算累积 mdat 负载大小和各文件原始基址 - cumPayload := make([]uint64, len(files)) - totalPayload := uint64(0) - origBase := make([]uint64, len(files)) - for i, f := range files { - cumPayload[i] = totalPayload - totalPayload += uint64(len(f.Mdat.Data)) - origBase[i] = f.Ftyp.Size() + f.Moov.Size() + f.Mdat.HeaderSize() + if err := os.WriteFile(filelist, []byte(strings.Join(lines, "\n")+"\n"), 0644); err != nil { + return fmt.Errorf("创建文件列表失败: %w", err) } + defer os.Remove(filelist) - // 3. 以第一个文件为蓝本构建输出结构 - outFile := mp4.NewFile() - outFile.Ftyp = files[0].Ftyp - outFile.Moov = files[0].Moov - - outFile.Mdat = &mp4.MdatBox{} - outFile.Mdat.SetData(make([]byte, totalPayload)) - pos := 0 - for _, f := range files { - copy(outFile.Mdat.Data[pos:], f.Mdat.Data) - pos += len(f.Mdat.Data) - } - - // 4. 按轨道类型初始化合并状态 - type mergeState struct { - stbl *mp4.StblBox - chunkBase uint32 // 当前已合并的 chunk 数量(stsc 偏移用) - } - - states := make(map[string]*mergeState) - for _, trak := range files[0].Moov.Traks { - hdlr := trak.Mdia.Hdlr.HandlerType - if _, ok := states[hdlr]; ok { - continue - } - stbl := trak.Mdia.Minf.Stbl - nChunks := chunkCount(stbl) - states[hdlr] = &mergeState{stbl: stbl, chunkBase: nChunks} - } - - // 5. 合并后续文件 - for fi := 1; fi < len(files); fi++ { - f := files[fi] - for _, srcTrak := range f.Moov.Traks { - hdlr := srcTrak.Mdia.Hdlr.HandlerType - st, ok := states[hdlr] - if !ok { - continue - } - srcStbl := srcTrak.Mdia.Minf.Stbl - nChunks := chunkCount(srcStbl) - - // stco/co64:偏移调整 = 第一个文件的基址差 + 累积 mdat 偏移 - adjust := int64(origBase[0]-origBase[fi]) + int64(cumPayload[fi]) - mergeChunkOffsets(st.stbl, srcStbl, adjust) - - // stsz:合并采样大小表 - if err := mergeSampleSizes(st.stbl.Stsz, srcStbl.Stsz); err != nil { - return err - } - - // stts:合并时域采样表 - mergeTimeToSample(st.stbl.Stts, srcStbl.Stts) - - // stsc:合并样块映射表,FirstChunk 需平移 - if srcStbl.Stsc != nil { - for i, e := range srcStbl.Stsc.Entries { - sid := srcStbl.Stsc.GetSampleDescriptionID(i + 1) - if err := st.stbl.Stsc.AddEntry(e.FirstChunk+st.chunkBase, e.SamplesPerChunk, sid); err != nil { - return fmt.Errorf("stsc 合并失败: %w", err) - } - } - } - - st.chunkBase += nChunks - } - } - - // 6. 修正所有 chunk 偏移(合并后 moov 大小变化) - combinedBase := outFile.Ftyp.Size() + outFile.Moov.Size() + outFile.Mdat.HeaderSize() - finalAdjust := int64(combinedBase - origBase[0]) - for _, st := range states { - adjustChunkOffsets(st.stbl, finalAdjust) - } - - // 7. 更新 mvhd 总时长 - totalDur := uint64(0) - for _, f := range files { - totalDur += f.Moov.Mvhd.Duration - } - outFile.Moov.Mvhd.Duration = totalDur - - // 更新每个 tkhd 时长 - for _, trak := range outFile.Moov.Traks { - hdlr := trak.Mdia.Hdlr.HandlerType - totalDur = 0 - for _, f := range files { - for _, t := range f.Moov.Traks { - if t.Mdia.Hdlr.HandlerType == hdlr { - totalDur += t.Tkhd.Duration - break - } - } - } - trak.Tkhd.Duration = totalDur - } - - // 8. 重建顶级 Children(确保 mdat 指向合并后的数据) - outFile.Children = []mp4.Box{outFile.Ftyp, outFile.Moov, outFile.Mdat} - - // 9. 写入输出文件 - return mp4.WriteToFile(outFile, output) -} - -// chunkCount 返回 stbl 中的 chunk 数量 -func chunkCount(stbl *mp4.StblBox) uint32 { - if stbl.Stco != nil { - return uint32(len(stbl.Stco.ChunkOffset)) - } - if stbl.Co64 != nil { - return uint32(len(stbl.Co64.ChunkOffset)) - } - return 0 -} - -// mergeChunkOffsets 将 src 的 stco/co64 追加到 base,每个偏移加上 adjust -func mergeChunkOffsets(base, src *mp4.StblBox, adjust int64) { - if src.Stco != nil { - if base.Stco == nil { - base.Stco = &mp4.StcoBox{} - } - for _, off := range src.Stco.ChunkOffset { - base.Stco.ChunkOffset = append(base.Stco.ChunkOffset, uint32(int64(off)+adjust)) - } - } - if src.Co64 != nil { - if base.Co64 == nil { - base.Co64 = &mp4.Co64Box{} - } - for _, off := range src.Co64.ChunkOffset { - base.Co64.ChunkOffset = append(base.Co64.ChunkOffset, uint64(int64(off)+adjust)) - } - } -} - -// adjustChunkOffsets 对 stbl 中所有 stco/co64 偏移统一加上 adjust -func adjustChunkOffsets(stbl *mp4.StblBox, adjust int64) { - if stbl.Stco != nil { - for i := range stbl.Stco.ChunkOffset { - stbl.Stco.ChunkOffset[i] = uint32(int64(stbl.Stco.ChunkOffset[i]) + adjust) - } - } - if stbl.Co64 != nil { - for i := range stbl.Co64.ChunkOffset { - stbl.Co64.ChunkOffset[i] = uint64(int64(stbl.Co64.ChunkOffset[i]) + adjust) - } - } -} - -// mergeSampleSizes 合并 stsz 采样大小表 -func mergeSampleSizes(base, src *mp4.StszBox) error { - if src == nil { - return nil - } - if base.SampleUniformSize > 0 { - if src.SampleUniformSize == base.SampleUniformSize { - base.SampleNumber += src.SampleNumber - } else { - sizes := make([]uint32, base.SampleNumber) - for i := range sizes { - sizes[i] = base.SampleUniformSize - } - if src.SampleUniformSize > 0 { - for i := uint32(0); i < src.SampleNumber; i++ { - sizes = append(sizes, src.SampleUniformSize) - } - } else { - sizes = append(sizes, src.SampleSize...) - } - base.SampleUniformSize = 0 - base.SampleSize = sizes - base.SampleNumber = uint32(len(sizes)) - } - } else { - if src.SampleUniformSize > 0 { - for i := uint32(0); i < src.SampleNumber; i++ { - base.SampleSize = append(base.SampleSize, src.SampleUniformSize) - } - } else { - base.SampleSize = append(base.SampleSize, src.SampleSize...) - } - base.SampleNumber = uint32(len(base.SampleSize)) + cmd := exec.Command("ffmpeg", + "-f", "concat", + "-safe", "0", + "-i", filelist, + "-c", "copy", + "-y", output, + ) + if out, err := cmd.CombinedOutput(); err != nil { + return fmt.Errorf("ffmpeg concat 失败: %w, 输出: %s", err, string(out)) } return nil } -// mergeTimeToSample 合并 stts 时域采样表(相邻相同 delta 的条目合并) -func mergeTimeToSample(base, src *mp4.SttsBox) { - if src == nil || len(src.SampleCount) == 0 { - return - } - if len(base.SampleCount) > 0 { - lastDelta := base.SampleTimeDelta[len(base.SampleTimeDelta)-1] - firstDelta := src.SampleTimeDelta[0] - if lastDelta == firstDelta { - base.SampleCount[len(base.SampleCount)-1] += src.SampleCount[0] - base.SampleCount = append(base.SampleCount, src.SampleCount[1:]...) - base.SampleTimeDelta = append(base.SampleTimeDelta, src.SampleTimeDelta[1:]...) - return - } - } - base.SampleCount = append(base.SampleCount, src.SampleCount...) - base.SampleTimeDelta = append(base.SampleTimeDelta, src.SampleTimeDelta...) -} - func (s *dramaService) downloadFile(ctx context.Context, url, dest string) error { resp, err := http.Get(url) if err != nil { @@ -1701,12 +1921,9 @@ func (s *dramaService) downloadFile(ctx context.Context, url, dest string) error } func (s *dramaService) probeVideo(path string) error { - f, err := mp4.ReadMP4File(path) - if err != nil { - return fmt.Errorf("视频文件无效: %w", err) - } - if f.Moov == nil { - return fmt.Errorf("视频文件缺少 moov box") + cmd := exec.Command("ffprobe", "-v", "error", "-show_entries", "format=duration", "-of", "csv=p=0", path) + if out, err := cmd.CombinedOutput(); err != nil { + return fmt.Errorf("视频文件无效: %w, ffprobe 输出: %s", err, string(out)) } return nil } @@ -1751,14 +1968,19 @@ func cleanupEpisodeWorkspace(ctx context.Context, dramaTitle string, epIndex int } func calcSegDurs(episodeDuration int64, cfg *entity.ModelConfig) []int { - effectiveMax := cfg.MaxSingleDuration - if effectiveMax <= 0 { - effectiveMax = 30 - } - - minSingle := cfg.MinSingleDuration - if minSingle <= 0 { - minSingle = effectiveMax + // 从 video_schema.duration 读取模型单段时长约束 + effectiveMax := 15 // 默认值 + minSingle := 5 + if cfg.VideoSchema != "" { + var vs map[string]any + if err := json.Unmarshal([]byte(cfg.VideoSchema), &vs); err == nil { + if v := intVal(nested(vs, "parameters", "duration", "max"), 0); v > 0 { + effectiveMax = v + } + if v := intVal(nested(vs, "parameters", "duration", "min"), 0); v > 0 { + minSingle = v + } + } } if minSingle > effectiveMax { minSingle = effectiveMax @@ -1856,6 +2078,17 @@ func imageFileToBase64(path string) (string, error) { return "data:" + mime + ";base64," + base64.StdEncoding.EncodeToString(data), nil } +// parseMMSSToSeconds 将 "MM:SS" 格式的时间字符串转为总秒数 +func parseMMSSToSeconds(timeStr string) int { + parts := strings.Split(timeStr, ":") + if len(parts) == 2 { + m, _ := strconv.Atoi(parts[0]) + sec, _ := strconv.Atoi(parts[1]) + return m*60 + sec + } + return 0 +} + // ValidateUploadFile 校验上传文件的大小和格式是否符合模型配置约束 // 参数 fileSize 为文件字节数,fileName 为原始文件名 // 返回 error 表示文件不合规 diff --git a/shortdrama/service/episode_service.go b/shortdrama/service/episode_service.go index c7c3e87..c9dd201 100644 --- a/shortdrama/service/episode_service.go +++ b/shortdrama/service/episode_service.go @@ -6,6 +6,7 @@ import ( "fmt" "os" "path/filepath" + "sort" "strings" "time" @@ -50,18 +51,35 @@ func (s *dramaService) AddEpisode(ctx context.Context, dramaId int64, title, des } index = maxIdx + 1 } - id, err := dao.Episode.Insert(ctx, &entity.Episode{ - DramaId: dramaId, - Index: index, - Title: title, - Description: description, - Script: script, - Status: consts.EpisodeStatusPending, - }) - if err != nil { - return 0, err + // 无脚本时直接插入剧集(无需创建任务) + if script == "" { + return dao.Episode.Insert(ctx, &entity.Episode{ + DramaId: dramaId, Index: index, Title: title, + Description: description, Script: script, + Status: consts.EpisodeStatusPending, + }) } - return id, nil + // 有脚本时用事务同时写入剧集和任务 + var epId int64 + now := time.Now().Format("2006-01-02 15:04:05") + err := g.DB().Transaction(ctx, func(ctx context.Context, tx gdb.TX) error { + r, e := tx.Model(public.TableNameEpisode).Ctx(ctx).Data(g.Map{ + "drama_id": dramaId, "idx": index, "title": title, + "description": description, "script": script, + "status": consts.EpisodeStatusPending, + "created_at": now, "updated_at": now, + }).Insert() + if e != nil { + return e + } + id, e := r.LastInsertId() + if e != nil { + return e + } + epId = id + return createPendingTasks(ctx, dramaId, epId, script, tx) + }) + return epId, err } func (s *dramaService) UpdateEpisode(ctx context.Context, dramaId, epId int64, title, description, script string, index int) error { @@ -80,17 +98,29 @@ func (s *dramaService) UpdateEpisode(ctx context.Context, dramaId, epId int64, t } if script != "" { e.Script = script - if e.Status == "" { - e.Status = consts.EpisodeStatusPending - } + e.Status = consts.EpisodeStatusPending } if index > 0 { e.Index = index } - if err := dao.Episode.Update(ctx, epId, e); err != nil { - return err - } - return nil + + needsTaskUpdate := script != "" + return g.DB().Transaction(ctx, func(ctx context.Context, tx gdb.TX) error { + // 更新剧集 + if _, e := tx.Model(public.TableNameEpisode).Ctx(ctx).Data(e).Where("id", epId).Update(); e != nil { + return e + } + // 脚本变动时重创建任务 + if needsTaskUpdate { + if _, e := tx.Model(public.TableNameGenerationTask).Ctx(ctx).Where("episode_id", epId).Delete(); e != nil { + return e + } + if e := createPendingTasks(ctx, dramaId, epId, script, tx); e != nil { + return fmt.Errorf("重新创建待生成任务失败: %w", e) + } + } + return nil + }) } func (s *dramaService) DeleteEpisode(ctx context.Context, dramaId, epId int64) error { @@ -257,7 +287,322 @@ func (s *dramaService) buildScriptGenUserInput(d *entity.Drama, episodeTitle, de b.WriteString("\n") } - b.WriteString(fmt.Sprintf("【要求】\n请根据以上剧情描述,为当前剧集创作一份详细的剧本。剧情描述是核心创作依据,请严格围绕其内容展开。\n\n输出格式:JSON数组,每个元素是一个镜头对象(shot),包含以下字段:\n- index: 镜头序号(从1开始)\n- startTime: 开始时间(格式MM:SS)\n- endTime: 结束时间(格式MM:SS)\n- event: 事件描述(该镜头中发生的情节概括)\n- cameraMovement: 运镜方式,从以下标准类型中选择一种:固定镜头、推、拉、摇、移、跟、升、降、旋转、晃动、航拍\n- characters: 出演人物数组,填写演员名称,如[\"张三\", \"李四\"],从「可用演员」中选择\n- scene: 场景名称,从「可用场景」中选择\n- props: 道具数组,填写道具名称,如[\"剑\", \"酒杯\"],从「可用道具」中选择\n\n直接输出JSON数组,不要markdown代码块标记,不要其他任何内容。\n\n要求:\n1. 剧本必须紧扣剧情描述展开\n2. 所有镜头时长之和应等于 %d 秒(每个镜头5-15秒不等)\n3. 优先使用提供的演员、场景和道具,如需新增请合理创作\n4. 每个镜头的characters、scene、props字段必须从提供的可用列表中选取名称,确保与参考素材一致\n5. 对话自然流畅,情节有起伏\n6. 镜头数量建议6-15个\n", d.EpisodeDuration)) + b.WriteString(fmt.Sprintf("【要求】\n请根据以上剧情描述,为当前剧集创作一份详细的剧本。剧情描述是核心创作依据,请严格围绕其内容展开。\n\n输出格式:JSON数组,每个元素是一个镜头对象(shot),包含以下字段:\n- index: 镜头序号(从1开始)\n- startTime: 开始时间(格式MM:SS)\n- endTime: 结束时间(格式MM:SS)\n- event: 事件描述(该镜头中发生的情节概括)\n- cameraMovement: 运镜方式,从以下标准类型中选择一种:固定镜头、推、拉、摇、移、跟、升、降、旋转、晃动、航拍\n- characters: 出演人物数组,填写演员名称,如[\"张三\", \"李四\"],从「可用演员」中选择\n- scene: 场景名称,从「可用场景」中选择\n- props: 道具名称数组,如[\"剑\", \"酒杯\"],从「可用道具」中选择\n\n直接输出JSON数组,不要markdown代码块标记,不要其他任何内容。\n\n要求:\n1. 剧本必须紧扣剧情描述展开\n2. 所有镜头时长之和应等于 %d 秒(每个镜头%d-%d秒不等)\n3. 优先使用提供的演员、场景和道具,如需新增请合理创作\n4. 每个镜头的characters、scene、props字段必须从提供的可用列表中选取名称,确保与参考素材一致\n5. 对话自然流畅,情节有起伏\n6. 镜头数量建议6-15个\n", d.EpisodeDuration, d.MinShotDuration, d.MaxShotDuration)) return b.String() } + +// createPendingTasks 为剧集创建待生成任务(按 calcSegDurs 拆段),每段脚本只包含本段对应部分 +func createPendingTasks(ctx context.Context, dramaId, epId int64, script string, tx gdb.TX) error { + d, err := dao.Drama.GetOne(ctx, dramaId) + if err != nil || d == nil { + return fmt.Errorf("短剧不存在: %d", dramaId) + } + modelCfg := ConfigService.Get(ctx) + + // 预加载参考素材(演员优先,场景其次,道具最后) + chars, _, _ := dao.Character.ListPageByDrama(ctx, dramaId, 1, -1) + scenes, _ := dao.Scene.ListByDrama(ctx, dramaId) + props, _ := dao.Prop.ListByDrama(ctx, dramaId) + + // 读取反向提示词 + negativePrompt := "" + if data, err := os.ReadFile("negative_prompt.md"); err == nil { + negativePrompt = strings.TrimSpace(string(data)) + } + + // 从 video_schema 读取参考素材上限和默认参数 + refMax := 5 + var vs map[string]any + if modelCfg.VideoSchema != "" { + json.Unmarshal([]byte(modelCfg.VideoSchema), &vs) + } + if vs != nil { + if v := nested(vs, "input", "reference_urls", "total_max"); v != nil { + if n, ok := v.(float64); ok { + refMax = int(n) + } + } + } + + // 从 video_schema 读取 prompt 最大字符数 + promptMaxChars := intVal(nested(vs, "input", "prompt", "max_chars"), 0) + + // 构建参考列表:演员 base64 → 场景图片 base64 → 道具图片 base64(与提交视频时顺序一致) + type namedRef struct { + name string + url string + } + var namedRefs []namedRef + for _, ch := range chars { + if len(namedRefs) >= refMax { + break + } + if ch.PortraitPath != "" { + namedRefs = append(namedRefs, namedRef{name: ch.Name, url: ch.PortraitPath}) + } + } + for _, sc := range scenes { + if len(namedRefs) >= refMax { + break + } + if sc.ImagePath != "" { + namedRefs = append(namedRefs, namedRef{name: sc.Name, url: sc.ImagePath}) + } + } + for _, p := range props { + if len(namedRefs) >= refMax { + break + } + if p.ImagePath != "" { + namedRefs = append(namedRefs, namedRef{name: p.Name, url: p.ImagePath}) + } + } + + // 读取拆段时长约束(与 calcSegDurs 逻辑一致) + effectiveMax := 15 + minSingle := 5 + if vs != nil { + if v := intVal(nested(vs, "parameters", "duration", "max"), 0); v > 0 { + effectiveMax = v + } + if v := intVal(nested(vs, "parameters", "duration", "min"), 0); v > 0 { + minSingle = v + } + } + if minSingle > effectiveMax { + minSingle = effectiveMax + } + + // 从 video_schema 读取默认 parameters(含 duration/audio/watermark 等) + params := map[string]any{} + if vs != nil { + if paramsDef := nested(vs, "parameters"); paramsDef != nil { + if pm, ok := paramsDef.(map[string]any); ok { + for k, v := range pm { + paramDef, ok := v.(map[string]any) + if !ok { + continue + } + // 如果有 supported_models 且不包含当前模型,跳过 + if sm, ok := paramDef["supported_models"]; ok { + if models, ok := sm.([]any); ok { + found := false + for _, m := range models { + if ms, ok := m.(string); ok && ms == modelCfg.VideoModelName { + found = true + break + } + } + if !found { + continue + } + } + } + if def, ok := paramDef["default"]; ok { + params[k] = def + } + } + } + } + // 尺寸映射:从 drama 的 Resolution/AspectRatio 查 video_schema + if sizes := nested(vs, "parameters", "size", "sizes"); sizes != nil { + if sm, ok := sizes.(map[string]any); ok { + if v := resolveSizeFromSchema(sm, d.Resolution, d.AspectRatio); v != "" { + params["size"] = v + } + } + } + } + + // 预构建名称替换信息(characterN 引用) + type _nameRep struct{ name, label string } + var _reps []_nameRep + if len(namedRefs) > 0 { + for j, nr := range namedRefs { + _reps = append(_reps, _nameRep{nr.name, fmt.Sprintf("character%d", j+1)}) + } + sort.Slice(_reps, func(a, b int) bool { return len(_reps[a].name) > len(_reps[b].name) }) + } + var _refStr string + if len(_reps) > 0 { + var _parts []string + for _, r := range _reps { + _parts = append(_parts, fmt.Sprintf("%s=%s", r.label, r.name)) + } + _refStr = fmt.Sprintf("。角色引用说明:%s", strings.Join(_parts, "、")) + } + _refLen := len([]rune(_refStr)) + + // buildFinalPrompt 执行名称替换并追加引用说明 + buildFinalPrompt := func(raw string) string { + for _, r := range _reps { + raw = strings.ReplaceAll(raw, r.name, r.label) + } + return raw + _refStr + } + + // refURLs 在所有段中相同 + refURLs := make([]string, len(namedRefs)) + for j, nr := range namedRefs { + refURLs[j] = nr.url + } + + // ============ 按脚本类型分组 ============ + type _taskGroup struct { + promptText string + dur int + } + var taskGroups []_taskGroup + + if domain.IsShotsJSON(script) { + // JSON 镜头数组:按镜头分组,累加连续镜头直到触及时长或字符上限 + var allShots []domain.Shot + if err := json.Unmarshal([]byte(script), &allShots); err != nil || len(allShots) == 0 { + return fmt.Errorf("解析镜头脚本失败: %w", err) + } + type _shotGroup struct { + shots []domain.Shot + totalDur int + } + var cur _shotGroup + for _, sh := range allShots { + shDur := parseMMSSToSeconds(sh.EndTime) - parseMMSSToSeconds(sh.StartTime) + if shDur <= 0 { + shDur = 1 + } + // 计算加入当前镜头后的候选组总时长和提示词长度 + candShots := make([]domain.Shot, len(cur.shots)+1) + copy(candShots, cur.shots) + candShots[len(cur.shots)] = sh + candDur := cur.totalDur + shDur + candPrompt := buildFinalPrompt(domain.ShotsToText(candShots)) + + exceedDur := len(cur.shots) > 0 && candDur > effectiveMax + exceedChars := promptMaxChars > 0 && len([]rune(candPrompt)) > promptMaxChars + + if exceedDur || exceedChars { + // 当前组已满,保存并开启新组 + taskGroups = append(taskGroups, _taskGroup{ + promptText: buildFinalPrompt(domain.ShotsToText(cur.shots)), + dur: cur.totalDur, + }) + cur = _shotGroup{shots: []domain.Shot{sh}, totalDur: shDur} + } else { + cur.shots = candShots + cur.totalDur = candDur + } + } + if len(cur.shots) > 0 { + taskGroups = append(taskGroups, _taskGroup{ + promptText: buildFinalPrompt(domain.ShotsToText(cur.shots)), + dur: cur.totalDur, + }) + } + } else { + // 纯文本脚本:按 calcSegDurs 拆段 + splitScriptForSegment 切分(保留原逻辑) + segDurs := calcSegDurs(d.EpisodeDuration, modelCfg) + segStartTimes := make([]int, len(segDurs)) + for i := 1; i < len(segDurs); i++ { + segStartTimes[i] = segStartTimes[i-1] + segDurs[i-1] + } + totalDur := d.EpisodeDuration + + if promptMaxChars > 0 { + totalRunes := len([]rune(script)) + if totalRunes > 0 && _refLen < promptMaxChars { + charsPerSeg := promptMaxChars - _refLen + maxDurByChars := charsPerSeg * int(totalDur) / totalRunes + if maxDurByChars > 0 && maxDurByChars < effectiveMax { + if maxDurByChars < minSingle { + maxDurByChars = minSingle + } + segDurs = calcSegmentDurations(int(totalDur), maxDurByChars, minSingle) + segStartTimes = make([]int, len(segDurs)) + for i := 1; i < len(segDurs); i++ { + segStartTimes[i] = segStartTimes[i-1] + segDurs[i-1] + } + } + } + } + + for i, segDur := range segDurs { + segScript := splitScriptForSegment(script, segStartTimes[i], segDur, int(totalDur)) + taskGroups = append(taskGroups, _taskGroup{ + promptText: buildFinalPrompt(segScript), + dur: segDur, + }) + } + } + + // ============ 插入任务 ============ + for i, tg := range taskGroups { + segParams := make(map[string]any, len(params)+1) + for k, v := range params { + segParams[k] = v + } + segParams["duration"] = tg.dur + + bodyJSON, _ := json.Marshal(map[string]any{ + "model": modelCfg.VideoModelName, + "input": map[string]any{ + "prompt": tg.promptText, + "negative_prompt": negativePrompt, + "reference_urls": refURLs, + }, + "parameters": segParams, + }) + + _, err := tx.Model(public.TableNameGenerationTask).Ctx(ctx).Data(g.Map{ + "drama_id": dramaId, + "episode_id": epId, + "segment_idx": i, + "status": consts.TaskStatusPending, + "script": string(bodyJSON), + "duration": tg.dur, + "num_segments": len(taskGroups), + }).Insert() + if err != nil { + return fmt.Errorf("创建第%d段待生成任务失败: %w", i+1, err) + } + } + return nil +} + +// splitScriptForSegment 将整集脚本按时间范围切分为对应段落的脚本内容 +func splitScriptForSegment(fullScript string, segStartTime, segDur, totalDur int) string { + if domain.IsShotsJSON(fullScript) { + var shots []domain.Shot + if err := json.Unmarshal([]byte(fullScript), &shots); err != nil { + return fullScript + } + segEndTime := segStartTime + segDur + var filtered []domain.Shot + for _, s := range shots { + shStart := parseMMSSToSeconds(s.StartTime) + shEnd := parseMMSSToSeconds(s.EndTime) + if shEnd <= segStartTime || shStart >= segEndTime { + continue + } + filtered = append(filtered, s) + } + if len(filtered) > 0 { + return domain.ShotsToText(filtered) + } + return "" + } + // 纯文本脚本:按时长比例切分字符 + runes := []rune(fullScript) + if len(runes) == 0 || totalDur <= 0 { + return fullScript + } + startRune := len(runes) * segStartTime / totalDur + endRune := len(runes) * (segStartTime + segDur) / totalDur + if startRune >= len(runes) { + return "" + } + if endRune > len(runes) { + endRune = len(runes) + } + return string(runes[startRune:endRune]) +} diff --git a/shortdrama/service/generation_context.go b/shortdrama/service/generation_context.go index 8c09ba3..688eeeb 100644 --- a/shortdrama/service/generation_context.go +++ b/shortdrama/service/generation_context.go @@ -20,12 +20,13 @@ type OrderedRef struct { Index int // 在 OrderedRefs 中的位置 } -// GenerationContext 一次生成会话的上下文,预加载当前短剧的演员/场景/道具数据 +// GenerationContext 一次生成会话的上下文,预加载当前短剧的演员/场景/道具/背景音乐数据 type GenerationContext struct { - Drama *entity.Drama - Characters []*entity.Character - Scenes []*entity.Scene - Props []*entity.Prop + Drama *entity.Drama + Characters []*entity.Character + Scenes []*entity.Scene + Props []*entity.Prop + BackgroundMusic []*entity.BackgroundMusic // RefIndex 按名称索引,用于 Agent 输出匹配 RefIndex RefIndex @@ -51,12 +52,18 @@ func BuildGenerationContext(ctx context.Context, drama *entity.Drama) (*Generati return nil, fmt.Errorf("加载道具失败: %w", err) } + bgmList, err := dao.BackgroundMusic.ListByDrama(ctx, drama.Id) + if err != nil { + return nil, fmt.Errorf("加载背景音失败: %w", err) + } + ctx2 := &GenerationContext{ - Drama: drama, - Characters: characters, - Scenes: scenes, - Props: props, - RefIndex: make(RefIndex), + Drama: drama, + Characters: characters, + Scenes: scenes, + Props: props, + BackgroundMusic: bgmList, + RefIndex: make(RefIndex), } // 构建 RefIndex diff --git a/test.json b/test.json new file mode 100644 index 0000000..076c728 --- /dev/null +++ b/test.json @@ -0,0 +1 @@ +{"model": "wan2.6-r2v", "input": {"prompt": {"required": true, "type": "string", "max_chars": 1500, "description": "文本提示词。支持中英文,每个汉字/字母/标点占一个字符,超过部分自动截断。通过 character1、character2 引用参考角色,每个参考仅包含单一角色。模型仅通过此方式识别参考中的角色"}, "negative_prompt": {"required": false, "type": "string", "max_chars": 500, "description": "反向提示词,用来描述不希望在视频画面中出现的内容,可以对视频画面进行限制"}, "reference_urls": {"required": true, "type": "array", "items": "string", "ref_ordering": "按数组顺序定义角色顺序,第1个URL对应character1,第2个对应character2,依此类推。每个参考文件仅包含一个主体角色", "total_max": 5, "image_max": 5, "video_max": 3, "video_formats": ["MP4", "MOV"], "video_duration_sec": {"min": 1, "max": 30}, "video_max_size_mb": 100, "image_formats": ["JPEG", "JPG", "PNG", "BMP", "WEBP"], "image_note": "PNG不支持透明通道", "image_resolution_px": {"min": 240, "max": 8000}, "image_max_size_mb": 20}}, "parameters": {"size": {"required": false, "type": "string", "format": "{width}*{height}", "default": "1920*1080", "description": "视频分辨率,格式为宽*高。必须设置为具体数值(如1280*720),而不是1:1或720P", "sizes": {"720P": {"9:16": "720*1280", "16:9": "1280*720", "1:1": "960*960", "4:3": "1088*832", "3:4": "832*1088"}, "1080P": {"9:16": "1080*1920", "16:9": "1920*1080", "1:1": "1440*1440", "4:3": "1632*1248", "3:4": "1248*1632"}}}, "duration": {"required": false, "type": "integer", "min": 2, "max": 10, "default": 5, "description": "生成视频的时长,单位为秒"}, "shot_type": {"required": false, "type": "string", "enum": ["single", "multi"], "default": "single", "description": "镜头类型。参数优先级:shot_type > prompt。single=单镜头视频,multi=多镜头视频"}, "audio": {"required": false, "type": "boolean", "default": true, "supported_models": ["wan2.6-r2v-flash"], "description": "是否生成有声视频"}, "watermark": {"required": false, "type": "boolean", "default": false}, "seed": {"required": false, "type": "integer", "min": 0, "max": 2147483647}}} \ No newline at end of file diff --git a/tools/split_merged/split_merged.go b/tools/split_merged/split_merged.go new file mode 100644 index 0000000..f371f47 --- /dev/null +++ b/tools/split_merged/split_merged.go @@ -0,0 +1,263 @@ +package main + +import ( + "fmt" + "os" + "path/filepath" + "strconv" + + "github.com/Eyevinn/mp4ff/mp4" +) + +func main() { + if len(os.Args) < 4 { + fmt.Println("用法: split_merged <输入文件> <输出目录> <段数>") + os.Exit(1) + } + inputPath := os.Args[1] + outDir := os.Args[2] + numSegments, err := strconv.Atoi(os.Args[3]) + if err != nil || numSegments <= 0 { + fmt.Printf("无效的段数: %s\n", os.Args[3]) + os.Exit(1) + } + + // 读取合并文件 + f, err := mp4.ReadMP4File(inputPath) + if err != nil { + fmt.Printf("读取文件失败: %v\n", err) + os.Exit(1) + } + + if err := os.MkdirAll(outDir, 0755); err != nil { + fmt.Printf("创建输出目录失败: %v\n", err) + os.Exit(1) + } + + // 收集每轨道按段划分的 sample 范围 + type trakInfo struct { + trak *mp4.TrakBox + samples uint32 // 总 sample 数 + perSeg uint32 // 每段 sample 数 + } + + var traks []trakInfo + for _, trak := range f.Moov.Traks { + stbl := trak.Mdia.Minf.Stbl + if stbl.Stsz != nil { + totalSamples := stbl.Stsz.SampleNumber + perSeg := totalSamples / uint32(numSegments) + traks = append(traks, trakInfo{ + trak: trak, + samples: totalSamples, + perSeg: perSeg, + }) + fmt.Printf("轨道 %s: 总 %d samples, 每段 %d\n", trak.Mdia.Hdlr.HandlerType, totalSamples, perSeg) + } + } + + if len(traks) == 0 { + fmt.Println("未找到可分割的轨道") + os.Exit(1) + } + + // 逐段切分 + for seg := 0; seg < numSegments; seg++ { + outPath := filepath.Join(outDir, fmt.Sprintf("seg_%d.mp4", seg)) + fmt.Printf("切分第 %d 段 -> %s\n", seg, outPath) + + outFile := mp4.NewFile() + outFile.Ftyp = f.Ftyp + outFile.Moov = f.Moov // 稍后会修改 + + // 为每轨道创建仅包含本段 samples 的 stbl + for _, ti := range traks { + stbl := ti.trak.Mdia.Minf.Stbl + startSample := seg * int(ti.perSeg) + endSample := startSample + int(ti.perSeg) + if seg == numSegments-1 { + endSample = int(ti.samples) // 最后一段拿剩余所有 samples + } + + fmt.Printf(" 轨道 %s: samples [%d, %d)\n", ti.trak.Mdia.Hdlr.HandlerType, startSample, endSample) + + // 拷贝 stsz entry(只保留本段范围) + if stbl.Stsz != nil { + newStsz := &mp4.StszBox{} + if stbl.Stsz.SampleUniformSize > 0 { + newStsz.SampleUniformSize = stbl.Stsz.SampleUniformSize + newStsz.SampleNumber = uint32(endSample - startSample) + } else { + newStsz.SampleSize = make([]uint32, endSample-startSample) + copy(newStsz.SampleSize, stbl.Stsz.SampleSize[startSample:endSample]) + newStsz.SampleNumber = uint32(len(newStsz.SampleSize)) + } + stbl.Stsz = newStsz + } + + // 计算本段 chunk 范围:通过 stsc 找到对应的 chunk + // 直接根据 perSeg chunks 计算。每段 chunks = totalChunks / numSegments + chunkCount := uint32(0) + if stbl.Stco != nil { + chunkCount = uint32(len(stbl.Stco.ChunkOffset)) + } else if stbl.Co64 != nil { + chunkCount = uint32(len(stbl.Co64.ChunkOffset)) + } + perSegChunks := chunkCount / uint32(numSegments) + startChunk := seg * int(perSegChunks) + endChunk := startChunk + int(perSegChunks) + if seg == numSegments-1 { + endChunk = int(chunkCount) + } + + // 重新构建 stco(只保留本段 chunk) + if stbl.Stco != nil { + oldChunks := stbl.Stco.ChunkOffset + stbl.Stco.ChunkOffset = oldChunks[startChunk:endChunk] + // 调整偏移量:减去 mdat 中本段数据的起始位置 + baseOffset := oldChunks[startChunk] + for i := range stbl.Stco.ChunkOffset { + stbl.Stco.ChunkOffset[i] -= baseOffset + } + } + if stbl.Co64 != nil { + oldChunks := stbl.Co64.ChunkOffset + stbl.Co64.ChunkOffset = oldChunks[startChunk:endChunk] + baseOffset := oldChunks[startChunk] + for i := range stbl.Co64.ChunkOffset { + stbl.Co64.ChunkOffset[i] -= baseOffset + } + } + + // 重新构建 stsc(只保留本段 chunk 的 entries) + if stbl.Stsc != nil { + var newEntries []mp4.StscEntry + for _, e := range stbl.Stsc.Entries { + if e.FirstChunk-1 >= uint32(startChunk) && e.FirstChunk-1 < uint32(endChunk) { + newEntries = append(newEntries, mp4.StscEntry{ + FirstChunk: e.FirstChunk - uint32(startChunk), + SamplesPerChunk: e.SamplesPerChunk, + SampleDescriptionIndex: e.SampleDescriptionIndex, + }) + } + } + stbl.Stsc.Entries = newEntries + } + + // 重设 stts(重建,只保留本段 samples 对应的时间) + if stbl.Stts != nil { + newStts := &mp4.SttsBox{} + currentSample := uint32(0) + for i := range stbl.Stts.SampleCount { + count := stbl.Stts.SampleCount[i] + delta := stbl.Stts.SampleTimeDelta[i] + nextSample := currentSample + count + + if nextSample <= uint32(startSample) { + currentSample = nextSample + continue + } + if currentSample >= uint32(endSample) { + break + } + + overlapStart := uint32(0) + if currentSample < uint32(startSample) { + overlapStart = uint32(startSample) - currentSample + } + overlapEnd := count + if nextSample > uint32(endSample) { + overlapEnd = uint32(endSample) - currentSample + } + + if overlapStart < overlapEnd { + newStts.SampleCount = append(newStts.SampleCount, overlapEnd-overlapStart) + newStts.SampleTimeDelta = append(newStts.SampleTimeDelta, delta) + } + currentSample = nextSample + } + stbl.Stts = newStts + } + + // 修正 stco 偏移:加上 ftyp+moov+mdat_header 偏移 + mdatDataStart := uint32(0) // 稍后在写入时修正 + _ = mdatDataStart + } + + // 提取本段 mdat 数据 + mdatStart := uint32(0) + if len(traks) > 0 && traks[0].trak.Mdia.Minf.Stbl.Stco != nil { + // 使用第一个轨道的第一个 chunk 偏移量作为数据起始位置 + mdatStart = traks[0].trak.Mdia.Minf.Stbl.Stco.ChunkOffset[0] + } + + // 重构 mdat:从合并文件提取本段数据 + outFile.Mdat = &mp4.MdatBox{} + mdatPayloadLen := uint64(0) + // 计算本段 mdat 数据长度 + for _, ti := range traks { + stbl := ti.trak.Mdia.Minf.Stbl + if stbl.Stsz != nil { + if stbl.Stsz.SampleUniformSize > 0 { + mdatPayloadLen += uint64(stbl.Stsz.SampleUniformSize) * uint64(stbl.Stsz.SampleNumber) + } else { + for _, s := range stbl.Stsz.SampleSize { + mdatPayloadLen += uint64(s) + } + } + } + } + + if mdatPayloadLen > 0 { + segData := make([]byte, mdatPayloadLen) + // 从原始 mdat 拷贝 + offset := uint64(mdatStart) + _ = offset + copyPos := uint64(0) + for _, ti := range traks { + stbl := ti.trak.Mdia.Minf.Stbl + if stbl.Stsz != nil { + if stbl.Stsz.SampleUniformSize > 0 { + dataLen := uint64(stbl.Stsz.SampleUniformSize) * uint64(stbl.Stsz.SampleNumber) + copy(segData[copyPos:], f.Mdat.Data[offset:offset+dataLen]) + copyPos += dataLen + offset += dataLen + } else { + for _, s := range stbl.Stsz.SampleSize { + copy(segData[copyPos:], f.Mdat.Data[offset:offset+uint64(s)]) + copyPos += uint64(s) + offset += uint64(s) + } + } + } + } + outFile.Mdat.SetData(segData) + } + + // 更新 mdhd/tkhd/mvhd 时长 + outFile.Moov.Mvhd = f.Moov.Mvhd // 拷贝原始 mvhd + // 简化处理:直接使用原始 mvhd 时长 + + // 重建 Children + outFile.Children = []mp4.Box{outFile.Ftyp, outFile.Moov, outFile.Mdat} + + // 修正 stco 偏移:调整到实际文件位置 + combinedBase := outFile.Ftyp.Size() + outFile.Moov.Size() + outFile.Mdat.HeaderSize() + for _, ti := range traks { + stbl := ti.trak.Mdia.Minf.Stbl + if stbl.Stco != nil { + for i := range stbl.Stco.ChunkOffset { + stbl.Stco.ChunkOffset[i] += uint32(combinedBase) + } + } + } + + if err := mp4.WriteToFile(outFile, outPath); err != nil { + fmt.Printf("写入 %s 失败: %v\n", outPath, err) + continue + } + fmt.Printf(" 完成: %s\n", outPath) + } + + fmt.Println("切分完成") +}