#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
Samsung Galaxy Watch 固件转换工具 (NetOdin 专用兼容补丁包生成器) By ZGQ Inc. t.me/ZGQinc
功能概述：
  解压所有内部镜像，剔除导致 NetOdin 报错的 'meta-data' 目录；
  解压全部 '.lz4' 压缩分区，还原为 tar；
  额外生成一份剔除 PIT 的安全版 CSC ，避免跨区刷机重写分区表导致 RQT Write Fail；
  使用 Linux POSIX 'ustar' 格式重新打包，生成可直接被 NetOdin 秒加载的固件。

依赖环境：
  Python 3.8+
  pip install lz4
  推荐系统已安装 7-Zip 或系统自带 tar ，Windows 10/11 自带 bsdtar 或 WSL
"""

import os
import sys
import shutil
import subprocess
import time
import hashlib
import stat
import zipfile
import argparse

try:
    import lz4.frame
except ImportError:
    print("[错误] 缺少 lz4 模块，请先在终端运行: pip install lz4")
    sys.exit(1)

def safe_remove(path):
    if os.path.exists(path):
        try:
            os.chmod(path, stat.S_IWRITE)
        except Exception:
            pass
        try:
            os.remove(path)
        except Exception as e:
            time.sleep(0.2)
            try:
                os.chmod(path, stat.S_IWRITE)
                os.remove(path)
            except Exception:
                print(f"  [警告] 无法删除临时文件: {path} ({e})")

def safe_rmtree(path):
    def on_exc(action, p, exc_info):
        try:
            os.chmod(p, stat.S_IWRITE)
            os.remove(p)
        except Exception:
            pass
    if os.path.exists(path):
        shutil.rmtree(path, onerror=on_exc)

def decompress_lz4(in_path, out_path):
    fname = os.path.basename(in_path)
    t0 = time.time()
    total = 0
    with open(in_path, 'rb') as raw_in:
        with lz4.frame.LZ4FrameDecompressor() as decompressor, open(out_path, 'wb') as f_out:
            while chunk := raw_in.read(32 * 1024 * 1024):
                decomp = decompressor.decompress(chunk)
                if decomp:
                    f_out.write(decomp)
                    total += len(decomp)
    safe_remove(in_path)
    dt = time.time() - t0
    mb = total / (1024 * 1024)
    print(f"    解压完成: {fname} -> {os.path.basename(out_path)} ({mb:.1f} MB, 耗时 {dt:.2f}s)")
    return total

def run_ustar_tar(work_dir, output_tar_path, file_list):
    try:
        cmd = ["tar", "-C", work_dir, "--format=ustar", "-cf", output_tar_path] + file_list
        res = subprocess.run(cmd, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL)
        if res.returncode == 0 and os.path.exists(output_tar_path) and os.path.getsize(output_tar_path) > 0:
            return
    except Exception:
        pass

    if sys.platform == "win32":
        try:
            wsl_work = work_dir.replace("\\", "/")
            if ":" in wsl_work:
                drive, path = wsl_work.split(":", 1)
                wsl_work = f"/mnt/{drive.lower()}{path}"
            wsl_out = output_tar_path.replace("\\", "/")
            if ":" in wsl_out:
                drive, path = wsl_out.split(":", 1)
                wsl_out = f"/mnt/{drive.lower()}{path}"

            args_str = " ".join([f'"{f}"' for f in file_list])
            cmd = f'cd "{wsl_work}" && tar -H ustar -cf "{wsl_out}" {args_str}'
            res = subprocess.run(["wsl", "bash", "-c", cmd], stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL)
            if res.returncode == 0 and os.path.exists(output_tar_path) and os.path.getsize(output_tar_path) > 0:
                return
        except Exception:
            pass

    import tarfile
    with tarfile.open(output_tar_path, mode="w", format=tarfile.USTAR_FORMAT) as t:
        for f in file_list:
            fpath = os.path.join(work_dir, f)
            ti = t.gettarinfo(fpath, arcname=f)
            ti.uid = 1000
            ti.gid = 1000
            ti.uname = "dpi"
            ti.gname = "dpi"
            ti.mode = 0o644
            with open(fpath, "rb") as fp:
                t.addfile(ti, fp)

def calculate_md5_and_save(tar_path):
    md5_path = tar_path + ".md5"
    safe_remove(md5_path)
    
    md5 = hashlib.md5()
    with open(tar_path, "rb") as f:
        while chunk := f.read(32 * 1024 * 1024):
            md5.update(chunk)
    md5_hex = md5.hexdigest()
    
    shutil.copyfile(tar_path, md5_path)
    with open(md5_path, "ab") as f:
        line = f"{md5_hex}  {os.path.basename(tar_path)}\n".encode("ascii")
        f.write(line)
    return md5_hex

def process_tar_member(raw_tar_path, out_dir, slot_name):
    t_start = time.time()
    raw_fname = os.path.basename(raw_tar_path)
    base_name = raw_fname.replace(".tar.md5", "").replace(".tar", "")
    out_tar_name = f"{slot_name}_{base_name}_NetOdin.tar"
    out_tar_path = os.path.join(out_dir, out_tar_name)
    
    print(f"\n  正在处理组件 [{slot_name}]: {raw_fname}")
    work_dir = os.path.join(out_dir, f"__temp_{slot_name}__")
    safe_rmtree(work_dir)
    os.makedirs(work_dir, exist_ok=True)
    
    extracted = False
    try:
        res = subprocess.run(["7z", "x", raw_tar_path, f"-o{work_dir}", "-y"], stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL)
        if res.returncode == 0:
            extracted = True
    except Exception:
        extracted = False
        
    if not extracted:
        import tarfile
        with tarfile.open(raw_tar_path, mode="r:*") as t:
            t.extractall(work_dir)
            
    for root, dirs, files in os.walk(work_dir):
        for d in dirs:
            try: os.chmod(os.path.join(root, d), stat.S_IWRITE)
            except Exception: pass
        for f in files:
            try: os.chmod(os.path.join(root, f), stat.S_IWRITE)
            except Exception: pass

    for root, dirs, files in os.walk(work_dir):
        for d in list(dirs):
            if d.lower() == "meta-data":
                meta_path = os.path.join(root, d)
                print(f"    [清理] 移除引发报错的目录: {d}")
                safe_rmtree(meta_path)
                dirs.remove(d)

    lz4_tasks = []
    for root, dirs, files in os.walk(work_dir):
        for f in files:
            if f.endswith(".lz4"):
                in_p = os.path.join(root, f)
                out_p = os.path.join(root, f[:-4])
                lz4_tasks.append((in_p, out_p))
                
    for in_p, out_p in lz4_tasks:
        decompress_lz4(in_p, out_p)
        
    final_files = sorted(os.listdir(work_dir))
    pit_files = [f for f in final_files if f.endswith(".pit")]
    non_pit_files = [f for f in final_files if not f.endswith(".pit")]
    final_files = pit_files + non_pit_files

    safe_remove(out_tar_path)
    run_ustar_tar(work_dir, out_tar_path, final_files)
    md5_hex = calculate_md5_and_save(out_tar_path)
    tar_size = os.path.getsize(out_tar_path) / (1024 * 1024)
    print(f"    打包成功: {out_tar_name} ({tar_size:.1f} MB, MD5: {md5_hex})")

    if slot_name == "CSC" and pit_files:
        nopit_name = f"{slot_name}_{base_name}_NetOdin_NoPIT.tar"
        nopit_path = os.path.join(out_dir, nopit_name)
        safe_remove(nopit_path)
        run_ustar_tar(work_dir, nopit_path, non_pit_files)
        nopit_md5 = calculate_md5_and_save(nopit_path)
        nopit_size = os.path.getsize(nopit_path) / (1024 * 1024)
        print(f"    [推荐] 额外生成无分区表安全版: {nopit_name} ({nopit_size:.1f} MB)")

    safe_rmtree(work_dir)
    dt = time.time() - t_start
    print(f"    该组件处理完成，总耗时 {dt:.1f}s")

def inspect_and_process_zip(zip_path, root_dir):
    fname = os.path.basename(zip_path)
    print(f"\n{'='*70}")
    print(f"发现压缩包: {fname}")
    print(f"{'='*70}")
    
    try:
        with zipfile.ZipFile(zip_path, 'r') as z:
            namelist = z.namelist()
    except Exception as e:
        print(f"  [跳过] 无法作为 ZIP 读取: {e}")
        return

    has_bl = any(n.startswith("BL_") and (".tar" in n) for n in namelist)
    has_ap = any(n.startswith("AP_") and (".tar" in n) for n in namelist)
    has_csc = any((n.startswith("CSC_") or n.startswith("HOME_CSC_")) and (".tar" in n) for n in namelist)
    
    if not (has_bl and has_ap):
        print(f"  [跳过] 未在包内检测到标准的三星 BL/AP 固件组件。")
        return

    print("  [确认] 成功识别为三星官方多件套固件包！开始解压原始 tar 文件...")
    
    base_folder_name = os.path.splitext(fname)[0]
    out_dir = os.path.join(root_dir, f"{base_folder_name}_NetOdin_Ready")
    os.makedirs(out_dir, exist_ok=True)
    
    temp_zip_extract = os.path.join(out_dir, "__temp_raw_tars__")
    safe_rmtree(temp_zip_extract)
    os.makedirs(temp_zip_extract, exist_ok=True)
    
    try:
        res = subprocess.run(["7z", "x", zip_path, f"-o{temp_zip_extract}", "-y"], stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL)
        if res.returncode != 0:
            raise Exception()
    except Exception:
        with zipfile.ZipFile(zip_path, 'r') as z:
            z.extractall(temp_zip_extract)
            
    extracted_files = os.listdir(temp_zip_extract)
    slots = ["BL", "CSC", "USERDATA", "AP"]
    
    for slot in slots:
        matched = [f for f in extracted_files if f.startswith(f"{slot}_") and (".tar" in f)]
        if matched:
            raw_tar = os.path.join(temp_zip_extract, matched[0])
            process_tar_member(raw_tar, out_dir, slot)

    safe_rmtree(temp_zip_extract)
    
    print(f"\n{'='*70}")
    print(f"处理完成！输出至：")
    print(f"📂 {out_dir}")

def main():
    parser = argparse.ArgumentParser(description="三星手表固件 NetOdin 兼容转换工具")
    parser.add_argument("-d", "--dir", default=".", help="指定扫描固件的目录")
    args = parser.parse_args()
    
    target_dir = os.path.abspath(args.dir)
    print(f"Galaxy Watch 固件转换器 (NetOdin)")
    print(f"正在扫描目录: {target_dir}")
    
    zip_files = [os.path.join(target_dir, f) for f in os.listdir(target_dir) if f.lower().endswith(".zip")]
    
    if not zip_files:
        print("未在当前目录下发现任何 .zip 压缩包。")
        return
        
    print(f"找到 {len(zip_files)} 个压缩文件，开始逐一校验解析...")
    for zpath in zip_files:
        inspect_and_process_zip(zpath, target_dir)

if __name__ == "__main__":
    main()
