# -*- coding: utf-8 -*-
r"""convert_parts.py - STEP 어셈블리를 '부품 단위'로 분리해 OBJ로 뽑는다 [CEO 2026-07-19]

통짜 변환(step_to_web.py)과 달리, STEP 안의 조립 구조(XCAF 최상위 컴포넌트)를
부품 하나하나로 나눠 p001.obj, p002.obj ... 로 저장한다.
각 부품 정점은 '조립 상태의 절대 좌표(world/mm)'로 저장되므로,
웹에서 모든 부품을 원점(0,0,0)에 그대로 로드하면 원래 조립 모양이 복원된다.

- 부품 이름: STEP 라벨에 실제 이름이 있으면 그걸, 없으면(NAUO 자동명 등) 순번(pNNN).
- offset_mm: 컴포넌트 배치 이동량(참고용). 좌표는 이미 world에 구워져 있으므로 웹은 재적용 금지.
- 좌표 규약: 기존 fab_models 131개 OBJ와 동일 — 네이티브 STEP mm, 축 변환 없음.

사용: python convert_parts.py "입력.stp" 모델명
      -> fab_models/parts/<모델명>/p001.obj ... + parts.json
"""
import sys, os, re, json, time

from OCP.STEPCAFControl import STEPCAFControl_Reader
from OCP.TDocStd import TDocStd_Document
from OCP.XCAFApp import XCAFApp_Application
from OCP.XCAFDoc import XCAFDoc_DocumentTool
from OCP.TCollection import TCollection_ExtendedString
from OCP.TDF import TDF_LabelSequence, TDF_Label
from OCP.TDataStd import TDataStd_Name
from OCP.BRepMesh import BRepMesh_IncrementalMesh
from OCP.BRep import BRep_Tool
from OCP.TopExp import TopExp_Explorer
from OCP.TopAbs import TopAbs_FACE, TopAbs_REVERSED
from OCP.TopoDS import TopoDS
from OCP.Bnd import Bnd_Box
from OCP.BRepBndLib import BRepBndLib
from OCP.gp import gp_Trsf

LIN_DEFLECTION = 0.5   # mm  (통짜 변환과 동일)
ANG_DEFLECTION = 0.3   # rad


def label_name(label):
    n = TDataStd_Name()
    if label.FindAttribute(TDataStd_Name.GetID_s(), n):
        try:
            return n.Get().ToExtString()
        except Exception:
            return None
    return None


def is_real_name(name):
    """실제 부품명인지, STEP 자동 생성명(NAUO123)/공백인지 판별."""
    if not name:
        return False
    s = name.strip()
    if not s:
        return False
    if re.fullmatch(r"NAUO\d+", s):
        return False
    # 비 ASCII(모지바케 위험) 이름은 신뢰하지 않음
    try:
        s.encode("ascii")
    except Exception:
        return False
    return True


def load_doc(path):
    app = XCAFApp_Application.GetApplication_s()
    doc = TDocStd_Document(TCollection_ExtendedString("d"))
    app.NewDocument(TCollection_ExtendedString("MDTV-XCAF"), doc)
    reader = STEPCAFControl_Reader()
    reader.SetNameMode(True)
    if not reader.ReadFile(path):
        raise RuntimeError("STEP 읽기 실패: " + path)
    reader.Transfer(doc)
    return doc


def mesh_to_obj(shape, obj_path, model_name):
    """shape(월드 좌표)를 삼각분할해 OBJ로 저장. (정점수, 면수, bbox) 반환."""
    from OCP.TopLoc import TopLoc_Location
    BRepMesh_IncrementalMesh(shape, LIN_DEFLECTION, False, ANG_DEFLECTION, True)
    verts = []
    faces = []
    exp = TopExp_Explorer(shape, TopAbs_FACE)
    while exp.More():
        face = TopoDS.Face_s(exp.Current())
        l = TopLoc_Location()
        tri = BRep_Tool.Triangulation_s(face, l)
        if tri is None:
            exp.Next()
            continue
        trsf = l.Transformation()
        base = len(verts)
        n = tri.NbNodes()
        for i in range(1, n + 1):
            p = tri.Node(i)
            p = p.Transformed(trsf)   # <-- 핵심: 노드에 위치변환 적용(안 하면 원점에 쌓임)
            verts.append((p.X(), p.Y(), p.Z()))
        reversed_face = (face.Orientation() == TopAbs_REVERSED)
        nt = tri.NbTriangles()
        for i in range(1, nt + 1):
            t = tri.Triangle(i)
            a, b, c = t.Get()
            if reversed_face:
                a, c = c, a
            faces.append((base + a, base + b, base + c))
        exp.Next()

    if not verts or not faces:
        return 0, 0, None

    # bbox (월드 좌표)
    xs = [v[0] for v in verts]; ys = [v[1] for v in verts]; zs = [v[2] for v in verts]
    bbox = [round(max(xs) - min(xs), 1), round(max(ys) - min(ys), 1), round(max(zs) - min(zs), 1)]

    with open(obj_path, "w", encoding="utf-8") as f:
        f.write("# convert_parts.py  model=%s  coords=world_absolute mm\n" % model_name)
        for v in verts:
            f.write("v %.6f %.6f %.6f\n" % v)
        for a, b, c in faces:
            f.write("f %d %d %d\n" % (a, b, c))
    return len(verts), len(faces), bbox


def component_offset(shape):
    """shape TopLoc의 이동량(mm) 반환 (참고용)."""
    trsf = shape.Location().Transformation()
    t = trsf.TranslationPart()
    return [round(t.X(), 1), round(t.Y(), 1), round(t.Z(), 1)]


def _label_realname(st, comp):
    """comp 라벨 -> 참조(원형) 라벨 순으로 실제 부품명 탐색. (이름or None) 반환."""
    nm = label_name(comp)
    if not is_real_name(nm):
        ref = TDF_Label()
        if st.GetReferredShape_s(comp, ref):
            rn = label_name(ref)
            if is_real_name(rn):
                nm = rn
    if is_real_name(nm):
        return re.sub(r":\d+$", "", nm.strip())  # 인스턴스 접미(:1) 정리
    return None


def convert(src, model_name, depth=1, max_parts=None):
    t0 = time.time()
    out_dir = os.path.join("fab_models", "parts", model_name)
    os.makedirs(out_dir, exist_ok=True)

    print("[로드] %s" % os.path.basename(src))
    doc = load_doc(src)
    st = XCAFDoc_DocumentTool.ShapeTool_s(doc.Main())
    free = TDF_LabelSequence()
    st.GetFreeShapes(free)
    if free.Length() == 0:
        raise RuntimeError("최상위 shape 없음")
    root = free.Value(1)
    comps = TDF_LabelSequence()
    st.GetComponents_s(root, comps)
    total = comps.Length()
    if total == 0:
        # 어셈블리가 아님 -> 통짜 하나로
        comp_labels = [root]
    else:
        comp_labels = [comps.Value(i) for i in range(1, total + 1)]
    print("[구조] 최상위 부품 %d개 (depth=%d)" % (len(comp_labels), depth))

    parts = []
    fails = []
    counter = {"n": 0}

    def emit(shape, name, name_from, path):
        counter["n"] += 1
        idx = counter["n"]
        pfile = "p%03d.obj" % idx
        try:
            offset = component_offset(shape)
            nv, nf, bbox = mesh_to_obj(shape, os.path.join(out_dir, pfile), model_name)
            if nv == 0:
                fails.append({"file": pfile, "reason": "empty mesh"})
                print("  [%3d] %-28s SKIP(빈 메쉬)" % (idx, name))
                return
            parts.append({
                "file": pfile, "name": name, "name_from": name_from, "path": path,
                "verts": nv, "faces": nf, "bbox_mm": bbox, "offset_mm": offset,
            })
            print("  [%3d] %-28s v=%-6d f=%-6d bbox=%s" % (idx, name, nv, nf, bbox))
        except Exception as e:
            fails.append({"file": pfile, "reason": str(e)[:120]})
            print("  [%3d] %s ERROR: %s" % (idx, pfile, str(e)[:100]))

    for tidx, comp in enumerate(comp_labels, start=1):
        # 최상위 컴포넌트의 '조립 위치가 적용된' shape (월드 좌표)
        top_shape = st.GetShape_s(comp)

        # depth 2: 하위 컴포넌트 2개 이상이면 분할
        n_sub = 0
        subcomps = TDF_LabelSequence()
        if depth >= 2:
            ref = TDF_Label()
            if st.GetReferredShape_s(comp, ref):
                st.GetComponents_s(ref, subcomps)
                n_sub = subcomps.Length()

        # 대분류(최상위 컴포넌트) 경로 라벨: 실명 없으면 그룹 순번 G%02d
        top_real = _label_realname(st, comp)
        top_path = top_real if top_real else ("G%02d" % tidx)

        if depth >= 2 and n_sub >= 2:
            # top 컴포넌트 shape의 Location(월드 배치)을 꺼내 하위 shape에 적용
            loc = top_shape.Location()
            for sidx in range(1, n_sub + 1):
                subcomp = subcomps.Value(sidx)
                sub_world = st.GetShape_s(subcomp).Moved(loc)
                sname = _label_realname(st, subcomp)
                if sname:
                    emit(sub_world, sname, "step", "%s/%s" % (top_path, sname))
                else:
                    emit(sub_world, "p%03d_%d" % (tidx, sidx), "seq",
                         "%s/G%02d-%d" % (top_path, tidx, sidx))
        else:
            # 리프/단품 -> 통짜 그대로 (depth1 동작과 동일)
            nm = _label_realname(st, comp)
            if nm:
                emit(top_shape, nm, "step", top_path)
            else:
                emit(top_shape, "p%03d" % tidx, "seq", top_path)

    if max_parts is not None and len(parts) > max_parts:
        print("[경고] 방출 부품 %d개 > max_parts %d (중단 안 함)" % (len(parts), max_parts))

    manifest = {
        "model": model_name,
        "src": src,
        "coords": "world_absolute",   # OBJ 정점이 이미 조립 절대좌표(mm). 웹은 offset 재적용 금지.
        "unit": "mm",
        "part_count": len(parts),
        "parts": parts,
    }
    if fails:
        manifest["failures"] = fails
    with open(os.path.join(out_dir, "parts.json"), "w", encoding="utf-8") as f:
        json.dump(manifest, f, ensure_ascii=False, indent=1)

    tot_mb = sum(os.path.getsize(os.path.join(out_dir, p["file"])) for p in parts) / 1024 / 1024
    print("[완료] %d부품, %.1fMB, 실패 %d개, %.0f초 -> %s" %
          (len(parts), tot_mb, len(fails), time.time() - t0, out_dir))
    return manifest


if __name__ == "__main__":
    import argparse
    ap = argparse.ArgumentParser(
        description="STEP 어셈블리를 부품 단위 OBJ로 분리")
    ap.add_argument("src", help="입력 STEP 경로")
    ap.add_argument("model_name", help="출력 모델(폴더)명")
    ap.add_argument("--depth", type=int, default=1,
                    help="1=최상위 컴포넌트 단위(기본), 2=하위 컴포넌트까지 분할")
    ap.add_argument("--max-parts", type=int, default=None,
                    help="방출 부품 수 상한(초과 시 경고만)")
    a = ap.parse_args()
    convert(a.src, a.model_name, depth=a.depth, max_parts=a.max_parts)
