#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
分批拼接双语视频，每批20个片段，然后合并所有批次
"""

import os
import subprocess
import tempfile


def get_video_files(video_dir):
    """
    获取所有视频文件列表
    """
    video_files = []
    for filename in os.listdir(video_dir):
        if filename.endswith('.mp4') and filename.startswith('video_'):
            idx = int(filename.split('_')[1].split('.')[0])
            video_files.append((idx, os.path.join(video_dir, filename)))
    
    video_files.sort(key=lambda x: x[0])
    return video_files


def concat_batch(video_pairs, output_path, batch_idx):
    """
    拼接一批视频
    """
    num_segments = len(video_pairs)
    
    # 构建输入参数
    input_args = []
    for ch_path, en_path in video_pairs:
        input_args.extend(['-i', ch_path])
        input_args.extend(['-i', en_path])
    
    # 构建filter_complex
    filter_parts = []
    
    # 处理每个输入
    for i in range(num_segments * 2):
        filter_parts.append(f'[{i}:v]scale=1080:1920:force_original_aspect_ratio=decrease,pad=1080:1920:(ow-iw)/2:(oh-ih)/2,setsar=1,fps=25[v{i}]')
        filter_parts.append(f'[{i}:a]aformat=sample_rates=44100:channel_layouts=stereo[a{i}]')
    
    # 拼接所有视频和音频
    video_inputs = ''.join([f'[v{i}]' for i in range(num_segments * 2)])
    audio_inputs = ''.join([f'[a{i}]' for i in range(num_segments * 2)])
    
    filter_parts.append(f'{video_inputs}concat=n={num_segments * 2}:v=1:a=0[outv]')
    filter_parts.append(f'{audio_inputs}concat=n={num_segments * 2}:v=0:a=1[outa]')
    
    filter_complex = ';'.join(filter_parts)
    
    # 构建完整命令
    cmd = [
        'ffmpeg', '-y'
    ] + input_args + [
        '-filter_complex', filter_complex,
        '-map', '[outv]',
        '-map', '[outa]',
        '-c:v', 'libx264',
        '-c:a', 'aac',
        '-b:a', '192k',
        output_path
    ]
    
    print(f"  正在拼接批次 {batch_idx + 1}...")
    result = subprocess.run(cmd, capture_output=True, text=True, encoding='utf-8', errors='ignore')
    
    return result.returncode == 0


def main():
    chinese_video_dir = "Artemis_II_十天绕月_videos"
    english_video_dir = "Artemis_II_十天绕月_tts_videos"
    output_file = "Artemis_II_十天绕月_bilingual.mp4"
    
    chinese_videos = get_video_files(chinese_video_dir)
    english_videos = get_video_files(english_video_dir)
    
    print(f"找到 {len(chinese_videos)} 个中文视频片段")
    print(f"找到 {len(english_videos)} 个英文TTS视频片段")
    
    if len(chinese_videos) != len(english_videos):
        print("错误：中文视频和英文TTS视频数量不匹配！")
        return
    
    # 构建视频对列表
    video_pairs = []
    for (ch_idx, ch_path), (en_idx, en_path) in zip(chinese_videos, english_videos):
        video_pairs.append((ch_path, en_path))
    
    # 分批处理，每批10个视频对（20个片段）
    batch_size = 10
    batches = [video_pairs[i:i+batch_size] for i in range(0, len(video_pairs), batch_size)]
    
    print(f"总共 {len(video_pairs)} 对视频，分为 {len(batches)} 批处理")
    
    # 创建临时目录存放批次文件
    temp_dir = tempfile.mkdtemp()
    batch_files = []
    
    try:
        # 处理每个批次
        for i, batch in enumerate(batches):
            batch_file = os.path.join(temp_dir, f"batch_{i:03d}.mp4")
            batch_files.append(batch_file)
            
            if not concat_batch(batch, batch_file, i):
                print(f"  批次 {i + 1} 拼接失败！")
                return
            
            print(f"  批次 {i + 1} 完成")
        
        # 合并所有批次
        print(f"\n正在合并 {len(batch_files)} 个批次...")
        
        # 创建concat文件列表
        concat_file = os.path.join(temp_dir, "concat_list.txt")
        with open(concat_file, 'w', encoding='utf-8') as f:
            for batch_file in batch_files:
                f.write(f"file '{batch_file}'\n")
        
        # 使用concat文件列表合并
        cmd = [
            'ffmpeg', '-y',
            '-f', 'concat',
            '-safe', '0',
            '-i', concat_file,
            '-c', 'copy',
            output_file
        ]
        
        result = subprocess.run(cmd, capture_output=True, text=True, encoding='utf-8', errors='ignore')
        
        if result.returncode == 0:
            print(f"成功！完整的双语视频保存到: {output_file}")
            if os.path.exists(output_file):
                file_size = os.path.getsize(output_file)
                print(f"文件大小: {file_size / 1024 / 1024:.2f} MB")
        else:
            print(f"错误：合并失败 - {result.stderr}")
    
    finally:
        # 清理临时文件
        import shutil
        if os.path.exists(temp_dir):
            shutil.rmtree(temp_dir)


if __name__ == "__main__":
    main()
