#!/bin/bash
set -e

# ============================================================
# Passport 生成环境部署脚本 - 最终版
# 目标机器：支持 opencode，联网环境
# 包含：m03 模板坐标 + 8 可编辑字段 + 4 固定字段 + MRZ
# ============================================================

echo "🚀 Deploying passport generation environment (final version)..."

# 1. 安装系统依赖
echo "📦 Installing system packages..."
apt-get update && apt-get install -y --no-install-recommends \
    python3 python3-pip python3-venv \
    fonts-dejavu-core \
    wget curl \
    && rm -rf /var/lib/apt/lists/*

# 2. 安装 Python 依赖
echo "🐍 Installing Python packages..."
pip3 install --break-system-packages --no-cache-dir \
    pillow==10.4.0 \
    opencv-python-headless==4.10.0.84 \
    numpy==1.26.4 \
    flask

# 3. 创建目录结构
echo "📁 Creating directories..."
mkdir -p /data/passport/{fonts,models,input,output,scripts,data_json}
mkdir -p /data/passport/fonts
mkdir -p /data/passport/models

# 4. 安装字体
echo "🔤 Installing fonts..."
cp /usr/share/fonts/truetype/dejavu/DejaVuSans*.ttf /data/passport/fonts/ 2>/dev/null || true

# 5. 下载 LaMa 模型
echo "📥 Downloading LaMa model..."
cd /data/passport/models
if [ ! -f big-lama.pt ]; then
    wget -q "https://github.com/Sanster/models/releases/download/add_big_lama/big-lama.pt" -O big-lama.pt
fi

# 6. 部署核心脚本
echo "📝 Deploying core scripts..."

# 6.1 主写入脚本
cat > /data/passport/scripts/passport_m03.py << 'PYEOF'
#!/usr/bin/env python3
"""
Passport Generation API - m03 Template Version
支持：8 个可编辑字段 + 4 个固定字段 + MRZ 自动生成
基于 m03 OCR 坐标校准的字段位置常量
"""

# ============================================================
# m03 模板字段位置常量 (OCR 校准)
# ============================================================

# 模板尺寸
M03_WIDTH = 1672
M03_HEIGHT = 1840

# Scale 因子
M03_SCALE_F = min(M03_WIDTH / 900, M03_HEIGHT / 630)  # 1.858

# 字段位置映射 (x, y) - 基于 m03 OCR 结果
# 格式: (x_label, y_label, x_value, y_value)
M03_FIELD_POS = {
    # 固定字段 (4个)
    "type":          (626, 1065, 626, 1065),      # PV - 左上
    "country_code":  (784, 1066, 784, 1066),      # MMR - 右上
    "nationality":   (600, 1200, 600, 1200),      # MYANMAR - 左列
    "authority":     (1091, 1397, 1091, 1397),    # MOHA, KYAINGTONG - 右列
    
    # 可编辑字段 (8个) - Web 页面变量
    "surname":       (802, 1039, 802, 1039),      # Surname / Nom - 右列第1行
    "given_name":    (1083, 1039, 1083, 1039),    # Given Name / Prenoms - 右列第1行
    "sex":           (632, 1105, 632, 1105),      # Sex / Sexe - 左列第2行
    "dob":           (649, 1170, 649, 1170),      # Date of birth - 左列第3行
    "birth_place":   (655, 1240, 655, 1240),      # Place of birth - 左列第4行
    "issue_date":    (618, 1305, 618, 1305),      # Date of issue - 左列第5行
    "expiry_date":   (661, 1371, 661, 1371),      # Date of expiry - 左列第6行
    "passport_no":   (1642, 46,   1642, 46),      # Passport No - 右上角
}

# MRZ 位置 (来自 m02 OCR，m03 无 MRZ 残留)
M03_MRZ_POS = {
    "prefix":    (45, 1110),  # "P<"
    "line1":     (836, 1126),  # MRZ Line 1
    "line2":     (836, 1170),  # MRZ Line 2
}

# 字体大小
M03_FONT_SIZES = {
    "title":   max(12, int(24 * M03_SCALE_F)),  # 44
    "pass":    max(10, int(17 * M03_SCALE_F)),  # 31
    "label":   max(8,  int(11 * M03_SCALE_F)),  # 20
    "val":     max(9,  int(13 * M03_SCALE_F)),  # 24
    "mrz":     max(10, int(18 * M03_SCALE_F)),  # 33
    "tiny":    max(6,  int(8  * M03_SCALE_F)),  # 14
}

# 颜色
M03_COLORS = {
    "title": "#1a1a1a", "pass": "#6B0000", "gold": "#C4A35A",
    "label": "#666", "val": "#111", "mrz": "#111",
}

# ============================================================
# MRZ 生成 (ICAO 9303)
# ============================================================

WEIGHTS = (7, 3, 1)

def to_val(ch: str) -> int:
    ch = ch.upper()
    if ch == "<": return 0
    if ch.isdigit(): return int(ch)
    return ord(ch) - ord("A") + 10

def check_digit(data: str) -> int:
    return sum(to_val(c) * WEIGHTS[i % 3] for i, c in enumerate(data)) % 10

def _padded(raw: str, n: int) -> str:
    return raw.upper().ljust(n, "<")[:n]

def build_line1(surname: str, given: str = "", type_code: str = "PV", country: str = "MMR") -> str:
    s = surname.upper().replace(" ", "<")
    g = given.upper().replace(" ", "<") if given else ""
    t = type_code[1] if len(type_code) > 1 and type_code[1] != "<" else "<"
    c = country.upper()[:3]
    id_part = s + "<<" + g if g else s
    return ("P" + t + c + id_part).ljust(44, "<")[:44]

def build_line2(passport_no: str, birth_date: str, expiry_date: str,
                sex: str = "M", issuing_state: str = "MMR", personal_no: str = "") -> str:
    p = _padded(passport_no, 9)
    b = _padded(birth_date, 6)
    e = _padded(expiry_date, 6)
    n = _padded(personal_no, 14)
    
    ps = f"{p}{check_digit(p)}"
    bs = f"{b}{check_digit(b)}"
    es = f"{e}{check_digit(e)}"
    ns = n + str(check_digit(n)) if personal_no.strip("<") else n + "<"
    
    l43 = ps + issuing_state + bs + sex + es + ns
    return l43 + str(check_digit(ps + bs + es + ns))

def gen_mrz(data: dict) -> tuple[str, str]:
    from datetime import datetime
    dob = datetime.strptime(data["dob"], "%d %b %Y")
    expiry = datetime.strptime(data["expiry_date"], "%d %b %Y")
    return build_line1(data["surname"], data["given_name"], data.get("type", "PV"), data.get("country_code", "MMR")), \
           build_line2(data["passport_no"], dob.strftime("%y%m%d"), expiry.strftime("%y%m%d"), data.get("sex", "M"), "MMR")

# ============================================================
# 主写入函数 (Mask-based rendering)
# ============================================================

from PIL import Image, ImageDraw, ImageFont
from pathlib import Path

def load_font(name: str, size: int) -> ImageFont.FreeTypeFont:
    font_dir = Path(__file__).parent.parent / "fonts"
    font_files = {
        "title": "DejaVuSans-Bold.ttf",
        "pass": "DejaVuSans-Bold.ttf",
        "label": "DejaVuSans.ttf",
        "val": "DejaVuSans-Bold.ttf",
        "mrz": "OCR-B.otf",
        "tiny": "DejaVuSans.ttf",
    }
    if name not in font_files:
        raise ValueError(f"Unknown font name: {name}")
    path = font_dir / font_files[name]
    if not path.exists():
        raise FileNotFoundError(f"Font not found: {path}")
    return ImageFont.truetype(str(path), size)

def hex_to_rgb(hex_color: str) -> tuple:
    h = hex_color.lstrip('#')
    if len(h) == 3: h = ''.join(c*2 for c in h)
    return tuple(int(h[i:i+2], 16) for i in (0, 2, 4))

def draw_text_mask(img: Image.Image, xy: tuple, text: str, font: ImageFont.FreeTypeFont,
                   fill: str, anchor: str = "lt", debug: bool = False, label: str = "") -> None:
    """使用 font.getmask() 绘制文本，绕过 Pillow 9.5.0 draw.text() bug"""
    mask = font.getmask(text, mode='L')
    mask_img = Image.frombytes('L', mask.size, bytes(mask))
    mw, mh = mask_img.size
    x, y = xy
    
    # 计算左上角
    anchors = {
        "lt": (x, y), "mt": (x - mw//2, y), "rt": (x - mw, y),
        "mm": (x - mw//2, y - mh//2), "lb": (x, y - mh),
        "mb": (x - mw//2, y - mh), "rb": (x - mw, y - mh)
    }
    tl_x, tl_y = anchors.get(anchor, (x, y))
    
    # 创建文本层并贴图
    text_layer = Image.new('RGB', (mw, mh), hex_to_rgb(fill) if isinstance(fill, str) else fill)
    img.paste(text_layer, (tl_x, tl_y), mask_img)
    
    if debug:
        draw = ImageDraw.Draw(img)
        draw.rectangle([tl_x, tl_y, tl_x + mw, tl_y + mh], outline="#FF0000", width=2)
        draw.ellipse([x-3, y-3, x+3, y+3], fill="#00FF00")
        debug_font = load_font("sans", 12)
        draw.text((tl_x, tl_y - 14), f"{label} {mw}x{mh}", fill="#FF0000", font=debug_font)

def write_passport(img_path: str, out_path: str, data: dict, debug: bool = False, skip_header: bool = False) -> None:
    img = Image.open(img_path).convert("RGB")
    W, H = img.size
    
    # 加载字体
    F = {k: load_font(k, M03_FONT_SIZES[k]) for k in ["title", "pass", "label", "val", "mrz", "tiny"]}
    C = M03_COLORS
    
    draw = ImageDraw.Draw(img)
    
    if not skip_header:
        # 1. 页眉
        draw_text_mask(img, (W//2, 55), "REPUBLIC OF THE UNION OF MYANMAR", F["title"], C["title"], "mt", debug, "title1")
        draw_text_mask(img, (W//2, 55 + 66), "P A S S P O R T", F["pass"], C["pass"], "mt", debug, "title2")
        
        # 2. 右上区域
        draw_text_mask(img, (W - 30, 46), f"Passport No  {data['passport_no']}", F["label"], "#1a1a1a", "rt", debug, "passport_no")
        draw_text_mask(img, (W - 30, 83), f"Type  {data.get('type','PV')}    Code  {data.get('country_code','MMR')}", F["tiny"], "#555", "rt", debug, "type_code")
    
    # 3. 固定字段 (4个) - 直接写入值
    # Type (PV)
    draw_text_mask(img, (626, 1065), data.get("type", "PV"), F["val"], C["val"], "lt", debug, "type_val")
    # Country Code (MMR)
    draw_text_mask(img, (784, 1066), data.get("country_code", "MMR"), F["val"], C["val"], "lt", debug, "country_val")
    # Nationality (MYANMAR)
    draw_text_mask(img, (600, 1200), data.get("nationality", "MYANMAR"), F["val"], C["val"], "lt", debug, "nationality_val")
    # Authority
    draw_text_mask(img, (1091, 1397), data.get("authority", "MOHA, KYAINGTONG"), F["val"], C["val"], "lt", debug, "authority_val")
    
    # 4. 可编辑字段 (8个) - Web 页面传入的变量
    editable = [
        ("surname", 802, 1039, "Surname"),
        ("given_name", 1083, 1039, "GivenName"),
        ("sex", 632, 1105, "Sex"),
        ("dob", 649, 1170, "DOB"),
        ("birth_place", 655, 1240, "BirthPlace"),
        ("issue_date", 618, 1305, "IssueDate"),
        ("expiry_date", 661, 1371, "ExpiryDate"),
        ("passport_no", 1642, 46, "PassportNo"),
    ]
    
    for key, x, y, label in editable:
        val = data.get(key, "")
        # passport_no 在右上角用 rt anchor
        anchor = "rt" if key == "passport_no" else "lt"
        draw_text_mask(img, (x, y), str(val), F["val"], C["val"], anchor, debug, f"field_{label}")
    
    # 5. MRZ 生成
    l1, l2 = gen_mrz(data)
    draw_text_mask(img, (45, 1110), "P<", F["mrz"], C["mrz"], "lt", debug, "mrz_prefix")
    draw_text_mask(img, (836, 1126), l1, F["mrz"], C["mrz"], "mt", debug, "mrz_l1")
    draw_text_mask(img, (836, 1170), l2, F["mrz"], C["mrz"], "mt", debug, "mrz_l2")
    
    img.save(out_path, dpi=(300, 300))
    print(f"✓ Generated: {out_path} ({W}x{H})")


# ============================================================
# CLI 入口
# ============================================================

if __name__ == "__main__":
    import argparse, json, sys
    
    p = argparse.ArgumentParser(description="Passport generator - m03 template")
    p.add_argument("--img", required=True)
    p.add_argument("--out", required=True)
    p.add_argument("--data", required=True, help="JSON or @file.json")
    p.add_argument("--debug", action="store_true")
    p.add_argument("--skip-header", action="store_true")
    args = p.parse_args()
    
    if args.data.startswith("@"):
        with open(args.data[1:]) as f:
            data = json.load(f)
    else:
        data = json.loads(args.data)
    
    # 必需字段：8 可编辑 + 4 固定
    required = ["surname", "given_name", "nationality", "dob", "sex",
                "birth_place", "issue_date", "expiry_date", "authority",
                "passport_no", "type", "country_code"]
    missing = [k for k in required if k not in data]
    if missing:
        sys.exit(f"Missing fields: {missing}")
    
    write_passport(args.img, args.out, data, args.debug, args.skip_header)
PYEOF
    chmod +x /data/passport/scripts/passport_m03.py

# 6.2 MRZ 辅助脚本
cat > /data/passport/scripts/mrz.py << 'PYEOF'
WEIGHTS=(7,3,1)
def to_val(ch):
    ch=ch.upper()
    if ch=="<": return 0
    if ch.isdigit(): return int(ch)
    return ord(ch)-ord("A")+10
def check_digit(data): return sum(to_val(c)*WEIGHTS[i%3] for i,c in enumerate(data))%10
def _padded(raw,n): return raw.upper().ljust(n,"<")[:n]
def build_line1(surname, given="", type_code="PV", country="MMR"):
    s=surname.upper().replace(" ","<"); g=given.upper().replace(" ","<") if given else ""
    t=type_code[1] if len(type_code)>1 and type_code[1]!="<" else "<"
    c=country.upper()[:3]; idp=s+"<<"+g if g else s
    return ("P"+t+c+idp).ljust(44,"<")[:44]
def build_line2(passport_no, birth_date, expiry_date, sex="M", issuing_state="MMR", personal_no=""):
    p=_padded(passport_no,9); b=_padded(birth_date,6); e=_padded(expiry_date,6); n=_padded(personal_no,14)
    ps=f"{p}{check_digit(p)}"; bs=f"{b}{check_digit(b)}"; es=f"{e}{check_digit(e)}"
    ns=n+check_digit(n) if personal_no.strip("<") else n+"<"
    l43=ps+issuing_state+bs+sex+es+ns
    return l43+str(check_digit(ps+bs+es+ns))
PYEOF

# 6.3 生成掩码脚本
cat > /data/passport/scripts/gen_mask.py << 'PYEOF'
#!/usr/bin/env python3
import argparse, cv2, numpy as np

def main():
    p = argparse.ArgumentParser()
    p.add_argument("--img", required=True); p.add_argument("--mask", required=True); p.add_argument("--template", default="m03")
    args = p.parse_args()
    img = cv2.imread(args.img); h,w = img.shape[:2]
    mask = np.zeros((h,w), np.uint8)
    if args.template=="m03":
        boxes = [[536,939,1109,968],[302,1023,557,1062],[606,1031,648,1054],[741,1029,863,1050],[1026,1029,1140,1050],[604,1050,648,1081],[741,1052,828,1081],[600,1092,664,1119],[602,1160,697,1181],[598,1184,784,1217],[597,1224,714,1256],[601,1289,641,1322],[911,1292,1033,1313],[598,1361,724,1382],[911,1359,995,1380],[908,1380,1275,1415],[598,1432,728,1453],[910,1428,1086,1453],[1006,1457,1041,1489]]
    else: raise ValueError(args.template)
    sf = lambda v: max(1, int(v*min(w/900,h/630)))
    for x1,y1,x2,y2 in boxes:
        cv2.rectangle(mask,(max(0,x1-8),max(0,y1-8)),(min(w,x2+8),min(h,y2+8)),255,-1)
    protect = [(sf(280),sf(118),sf(445),sf(318)),(0,0,w,sf(100)),(0,sf(1080),w,sf(1200)),(0,0,sf(20),h),(w-sf(20),0,w,h),(0,0,w,sf(20)),(0,h-sf(20),w,h)]
    for x1,y1,x2,y2 in protect: cv2.rectangle(mask,(x1,y1),(x2,y2),0,-1)
    mask = cv2.morphologyEx(mask, cv2.MORPH_CLOSE, np.ones((5,5),np.uint8))
    cv2.imwrite(args.mask, mask)
    print(f"✓ Mask: {args.mask}")

if __name__=="__main__": main()
PYEOF
chmod +x /data/passport/scripts/gen_mask.py

# 6.4 LaMa inpaint 脚本
cat > /data/passport/scripts/lama_inpaint.py << 'PYEOF'
#!/usr/bin/env python3
import argparse, torch, cv2, numpy as np
from pathlib import Path

def pad(x, m=8):
    h,w = x.shape[:2]
    nh,nw = ((h+m-1)//m)*m, ((w+m-1)//m)*m
    ph,pw = nh-h, nw-w
    if x.ndim==3: return np.pad(x,((0,ph),(0,pw),(0,0)),mode="reflect"),(ph,pw)
    return np.pad(x,((0,ph),(0,pw)),mode="reflect"),(ph,pw)

def main():
    p = argparse.ArgumentParser()
    p.add_argument("--img", required=True); p.add_argument("--mask", required=True); p.add_argument("--out", required=True)
    p.add_argument("--model", default="/data/passport/models/big-lama.pt")
    args = p.parse_args()

    model = torch.jit.load(args.model, map_location="cpu"); model.eval()
    img = cv2.imread(args.img); mask = cv2.imread(args.mask, 0)
    img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB); h,w = img.shape[:2]
    img_pad,_ = pad(img); mask_pad,_ = pad(mask)
    img_t = torch.from_numpy(img_pad).float()/255.0; img_t = img_t.permute(2,0,1).unsqueeze(0)
    mask_t = torch.from_numpy(mask_pad).float()/255.0; mask_t = mask_t.unsqueeze(0).unsqueeze(0)
    with torch.no_grad(): out = model(img_t, mask_t)
    out = out[0].permute(1,2,0).cpu().numpy()
    out = np.clip(out*255,0,255).astype(np.uint8)
    out = cv2.cvtColor(out, cv2.COLOR_RGB2BGR)[:h,:w]
    cv2.imwrite(args.out, out)
    print(f"✓ {args.out}")

if __name__=="__main__": main()
PYEOF
chmod +x /data/passport/scripts/lama_inpaint.py

# 7. 生成干净模板（如果尚不存在）
echo "🧹 Generating clean template if needed..."
cd /data/passport
if [ ! -f input/m03_clean.png ]; then
    echo "  Generating mask..."
    python3 scripts/gen_mask.py --img /passport/模板/m03.jpeg --mask input/m03_mask.png --template m03
    echo "  Running LaMa inpainting..."
    python3 scripts/lama_inpaint.py --img /passport/模板/m03.jpeg --mask input/m03_mask.png --out input/m03_clean.png
fi

# 8. 设置权限
echo "🔐 Setting permissions..."
chown -R www-data:www-data /data/passport
chmod -R 755 /data/passport/scripts

echo ""
echo "✅ Deployment complete!"
echo ""
echo "📁 目录结构:"
echo "  /data/passport/"
echo "  ├── fonts/           # DejaVu 字体"
echo "  ├── models/          # big-lama.pt"
echo "  ├── input/           # m03_clean.png (干净底图), m03_mask.png"
echo "  ├── output/          # 生成结果"
echo "  ├── scripts/         # 核心脚本:"
echo "  │   ├── passport_m03.py   # 主生成脚本"
echo "  │   ├── mrz.py            # MRZ 生成"
echo "  │   ├── gen_mask.py       # 掩码生成"
echo "  │   └── lama_inpaint.py   # LaMa 修复"
echo "  └── data_json/       # 批量数据目录"
echo ""
echo "🚀 使用示例:"
echo ""
echo "  # 1. 单次生成 (命令行)"
echo "  python3 /data/passport/scripts/passport_m03.py \\"
echo "    --img /data/passport/input/m03_clean.png \\"
echo "    --out /data/passport/output/test.png \\"
echo "    --data '{\"surname\":\"LI\",\"given_name\":\"WEI\",\"nationality\":\"MYANMAR\",\"dob\":\"15 JAN 1995\",\"sex\":\"M\",\"birth_place\":\"YANGON\",\"issue_date\":\"10 MAR 2023\",\"expiry_date\":\"09 MAR 2028\",\"authority\":\"MOHA, KYAINGTONG\",\"passport_no\":\"MK888888\",\"type\":\"PV\",\"country_code\":\"MMR\"}' \\"
echo "    --skip-header"
echo ""
echo "  # 2. 批量生成"
echo "  mkdir -p /data/passport/data_json"
echo "  echo '[{\"surname\":\"LI\",\"given_name\":\"WEI\",...}]' > /data/passport/data_json/batch.json"
echo "  python3 /data/passport/scripts/passport_m03.py \\"
echo "    --img /data/passport/input/m03_clean.png \\"
echo "    --out /data/passport/output/batch_001.png \\"
echo "    --data @/data/passport/data_json/batch.json"
echo ""
echo "  # 3. 查看帮助"
echo "  python3 /data/passport/scripts/passport_m03.py --help"
echo ""
echo "💡 提示: 生成的护照将保存在 /data/passport/output/ 目录"