#!/usr/bin/env python
# -*- coding: utf-8 -*-
"""연결 불가 11부품 진단 시험대.

왜 필요한가: 전체 해석 로그는 '연결 불가로 자동 제외 11개' 만 말한다. 왜 떨어졌는지
(표면이 멀다 / 씨앗이 안 생겼다 / 씨앗은 생겼는데 MPC 가 안 걸렸다 / 이웃이 빠져
연쇄로 빠졌다)는 숫자로 말하지 않는다. 그 네 가지를 실측으로 갈라 적는다.

  python _fea_연결진단.py req    → 머신 하중(6고정/-Z/3톤) 요청 JSON 을 만든다
  python _fea_연결진단.py run    → 그 요청으로 전체 해석 1회 + 진단 덤프 저장
  python _fea_연결진단.py table  → 덤프로 부품별 원인 표 (해석 재실행 없음)

머신 BC 는 fab_fea.js machineBC() 를 그대로 옮긴 것이다(지면접촉 zmin+5mm 부품 전부
고정, 최상단 3부품에 -Z 로 하중 1/3씩). 화면과 같은 조건이어야 비교가 성립한다.
"""
import json
import os
import pickle
import sys

import numpy as np

import fea_server as F

HERE = os.path.dirname(os.path.abspath(__file__))
REQ = os.path.join(F.WORK_ROOT, 'machine_request.json')
DUMP = os.path.join(F.WORK_ROOT, 'conn_diag.pkl')
OUT = os.path.join(HERE, '_fea_연결진단_결과.json')
MACHINE_N = 29400.0          # 3톤 (1톤 = 9,800N) — fab_fea.js MACHINE_LOADS 와 같은 값


def _parts_xyz(q):
    """요청의 positions + part_sizes → 부품별 (n,3) 좌표 배열."""
    P = np.asarray(q['positions'], float).reshape(-1, 3)
    out, off = [], 0
    for c in q['part_sizes']:
        rows = int(c) // 3
        out.append(P[off:off + rows])
        off += rows
    return out


def make_req():
    """화면 머신 프리셋과 같은 BC(6고정/-Z/3톤)를 붙인 요청을 만든다."""
    src = json.load(open(os.path.join(F.WORK_ROOT, 'last_request.json'), encoding='utf-8'))
    pts = _parts_xyz(src)
    box = [(p.min(axis=0), p.max(axis=0)) for p in pts]
    zmin = min(b[0][2] for b in box)
    fix, fixn = [], []
    for k, (lo, hi) in enumerate(box):
        if lo[2] <= zmin + 5:
            rr = max(30.0, 0.5 * min(hi[0] - lo[0], hi[1] - lo[1]))
            fix.append({'p': [float((lo[0] + hi[0]) / 2), float((lo[1] + hi[1]) / 2),
                              float(lo[2] + 3)], 'r': float(rr)})
            fixn.append(src['part_names'][k])
    tops = sorted(range(len(box)), key=lambda k: -box[k][1][2])[:3]
    load = []
    for k in tops:
        lo, hi = box[k]
        load.append({'p': [float((lo[0] + hi[0]) / 2), float((lo[1] + hi[1]) / 2),
                           float(hi[2] - 3)], 'r': 60.0, 'dir': [0, 0, -1],
                     'N': MACHINE_N / len(tops)})
    q = dict(src)
    q['fix_points'] = fix
    q['load_points'] = load
    q['load_N'] = MACHINE_N
    q['local_refine'] = 0
    json.dump(q, open(REQ, 'w', encoding='utf-8'))
    print('머신 요청 저장 %s' % REQ)
    print('  고정 %d개: %s' % (len(fix), ', '.join(fixn)))
    print('  하중 %d개(-Z, 각 %.0fN): %s'
          % (len(load), MACHINE_N / len(tops),
             ', '.join(src['part_names'][k] for k in tops)))
    return q


def run(local_n=0, diag=True, label='run'):
    q = json.load(open(REQ, encoding='utf-8'))
    if diag:
        os.environ['DZW_FEA_DIAG'] = DUMP
    else:
        os.environ.pop('DZW_FEA_DIAG', None)
    import time
    t0 = time.time()
    r = F.run_fea(q.get('positions') or [],
                  load_N=float(q.get('load_N')),
                  material=q.get('material', F.DEF_MAT),
                  E=q.get('E'), nu=q.get('nu'),
                  fix_points=q.get('fix_points'),
                  load_points=q.get('load_points'),
                  pick_radius=q.get('pick_radius'),
                  weld_mode=q.get('weld_mode') or 'bond',
                  weld_throat=float(q.get('weld_throat') or 0.0),
                  part_sizes=q.get('part_sizes'),
                  part_names=q.get('part_names'),
                  quality_bulk=True, weld_seed=True,
                  local_refine_n=local_n,
                  _logf=lambda m: print(m, flush=True))
    m, res = r['mesh'], r['result']
    ex = r.get('excluded_parts') or []
    rep = {'label': label, 'sec': round(time.time() - t0, 1),
           'nodes': m['nodes'], 'elems': m.get('elems'),
           'parts': m['parts'], 'welds': m.get('welds'),
           'excluded_n': len(ex), 'excluded': ex,
           'disp_mm': res.get('max_disp_mm'),
           'struct_max_MPa': res.get('max_vonmises_MPa'),
           'peak_MPa': res.get('peak_vonmises_MPa'),
           'basis': res.get('basis'),
           'bridges': (r['mesh'] or {}).get('weld_bridges'),
           'hull_parts': (r['mesh'] or {}).get('hull_parts'),
           'sheet_est': (r['mesh'] or {}).get('sheet_est'),
           'bc_warn': (r.get('bc') or {}).get('warn'),
           'candidates': len(r.get('candidates') or [])}
    print(json.dumps(rep, ensure_ascii=False, indent=1)[:3000])
    old = {}
    if os.path.exists(OUT):
        try:
            old = json.load(open(OUT, encoding='utf-8'))
        except Exception:
            old = {}
    old[label] = rep
    json.dump(old, open(OUT, 'w', encoding='utf-8'), ensure_ascii=False, indent=1)
    return r


# ── 표면 간격 실측 ──────────────────────────────────────────────────────────
def _min_gap(Vi, Fi, Vj, Fj, k=24):
    """부품 i 표본점 ↔ 부품 j 삼각형의 최소 거리(mm). 점↔삼각형 정확식.

    표본점은 꼭짓점 + 무게중심(weld_seeds_cad 과 같은 성격). 각 표본점마다 j 의
    삼각형 중 무게중심이 가까운 k개만 정확거리를 재고 최소를 취한다 — 전수는
    수만×수만이라 돌지 않는다. k 를 키우면 값은 단조 감소하므로 상한값이다.
    """
    from scipy.spatial import cKDTree
    Ti, Tj = Vi[Fi], Vj[Fj]
    S = np.vstack([Vi, Ti.mean(axis=1)])
    C = Tj.mean(axis=1)
    kk = min(k, len(C))
    _d, idx = cKDTree(C).query(S, k=kk)
    idx = np.atleast_2d(idx.T).T.reshape(len(S), kk)
    best = np.inf
    for col in range(kk):
        d = F.point_tri_dist_pairs(S, Tj[idx[:, col]])
        best = min(best, float(d.min()))
    return best


def table():
    d = pickle.load(open(DUMP, 'rb'))
    P = d['parts']
    no2i = {p['no']: i for i, p in enumerate(P)}
    grp, fixg = d['group'], set(d['fix_groups'])
    flo = d['floating']
    floset = set(flo)
    est = {e['no']: e for e in d['sheet_est']}
    lo = np.array([p['lo'] for p in P])
    hi = np.array([p['hi'] for p in P])
    # 부품쌍 → 씨앗/묶음 (자리번호 기준)
    seeds = {tuple(sorted(k)): v for k, v in d['seeds'].items()}
    ties = {tuple(sorted(k)): v for k, v in d['tie_pair'].items()}
    welds = {}
    for w in d['welds']:
        welds[tuple(sorted((w['i'], w['j'])))] = w
    rows = []
    for nn in flo:
        i = no2i[nn]
        Vi = np.asarray(P[i]['cadV'], float)
        Fi = np.asarray(P[i]['cadF'], np.int64)
        # 이웃 후보: 경계상자 거리가 가까운 30개(자기 제외)
        c = (lo + hi) / 2
        dd = np.maximum(np.maximum(lo - hi[i], lo[i] - hi), 0).max(axis=1)
        dd[i] = np.inf
        cand = np.argsort(dd)[:30]
        meas = []
        for j in cand:
            g = _min_gap(Vi, Fi, np.asarray(P[j]['cadV'], float),
                         np.asarray(P[j]['cadF'], np.int64))
            k = tuple(sorted((i, int(j))))
            meas.append({'no': P[j]['no'], 'name': P[j]['name'],
                         'gap_mm': round(g, 3),
                         'bbox_gap_mm': round(float(dd[j]), 3),
                         'bbox_overlap': bool(dd[j] <= 0.0),
                         'floating': P[j]['no'] in floset,
                         'in_fixgrp': grp[int(j)] in fixg,
                         'seed_n': seeds.get(k, 0),
                         'tie_n': ties.get(k, 0),
                         'weld_pair': k in welds,
                         'weld': welds.get(k)})
        meas.sort(key=lambda x: x['gap_mm'])
        fixed_side = [m for m in meas if m['in_fixgrp']]
        rows.append({
            'no': nn, 'name': P[i]['name'],
            'wall_mm': round(P[i]['wall_mm'], 2),
            'thickness_estimated': est.get(nn),
            'tie_total': sum(v for k, v in ties.items() if i in k),
            'seed_total': sum(v for k, v in seeds.items() if i in k),
            'weld_pairs': sum(1 for k in welds if i in k),
            'nearest': meas[:5],
            'nearest_fixed_side': fixed_side[:3],
            'min_gap_mm': meas[0]['gap_mm'] if meas else None,
            'min_gap_fixed_side_mm': fixed_side[0]['gap_mm'] if fixed_side else None,
        })
    out = {'wtol_mm': d['wtol_mm'], 'n_parts': len(P),
           'floating': flo, 'rows': rows,
           'failed': d['failed'], 'sheet_est': d['sheet_est']}
    p = os.path.join(HERE, '_fea_연결진단_표.json')
    json.dump(out, open(p, 'w', encoding='utf-8'), ensure_ascii=False, indent=1)
    print('tol=%.2fmm  부품 %d  떠있음 %d' % (d['wtol_mm'], len(P), len(flo)))
    print('%-5s %-26s %8s %8s %6s %6s %5s  %s'
          % ('번호', '이름', '최소간격', '고정쪽', '씨앗', '묶음', '추정', '가장 가까운 이웃'))
    for r in rows:
        n0 = r['nearest'][0] if r['nearest'] else {}
        print('%-5d %-26s %8.2f %8s %6d %6d %5s  %d번 %s%s'
              % (r['no'], r['name'][:26], r['min_gap_mm'] or -1,
                 ('%.2f' % r['min_gap_fixed_side_mm']) if r['min_gap_fixed_side_mm']
                 is not None else '-',
                 r['seed_total'], r['tie_total'],
                 ('%.1f' % r['thickness_estimated']['t_mm'])
                 if r['thickness_estimated'] else '-',
                 n0.get('no', -1), n0.get('name', '')[:20],
                 ' (떠있음)' if n0.get('floating') else ''))
    print('표 저장 %s' % p)
    return out


def replay(cand_mode='corner', margin=None, max_gap=None):
    """덤프(메시 포함)로 '다리 잇기' 만 다시 돌린다 — 전체 해석 없이 몇 초.

    cand_mode: 'corner'  = 겉면 모서리 절점만 종속 후보(surf_ids, 종전 동작)
               'midside' = 모서리 + 중간절점(faces6 전체) — C3D10 은 면마다
                           중간절점이 있어 후보 간격이 절반으로 촘촘해진다.
    """
    d = pickle.load(open(DUMP, 'rb'))
    if 'mesh' not in d or not d['mesh']:
        print('이 덤프에는 메시가 없다 — run 을 다시 돌려 메시 포함 덤프를 만들어라')
        return
    nodes = np.asarray(d['nodes'], float)
    P = d['parts']
    parts = []
    for i, p in enumerate(P):
        m = d['mesh'].get(i)
        q = {'part_no': p['no'], 'part_name': p['name'],
             'lo': np.asarray(p['lo'], float), 'hi': np.asarray(p['hi'], float),
             'cadV': np.asarray(p['cadV'], float), 'cadF': np.asarray(p['cadF'], np.int64)}
        if m:
            f6 = np.asarray(m['faces6'], np.int64)
            sid = (np.unique(f6) if cand_mode == 'midside'
                   else np.asarray(m['surf_ids'], np.int64))
            q.update({'faces6': f6, 'surf_ids': sid, 'surf_pts': nodes[sid],
                      'n0': m['n0'], 'n1': m['n1']})
        parts.append(q)
    if margin is not None:
        F.BRIDGE_TOL_MARGIN_MM = float(margin)
    if max_gap is None:
        max_gap = F.BRIDGE_MAX_GAP_MM
    # 떠 있는 부재와 그 이웃만 메시를 담았으니, 메시 없는 부품은 후보에서 뺀다
    ok = set(i for i, q in enumerate(parts) if 'faces6' in q)
    #   메시를 안 담은 부품은 '고정 성분' 에 넣어 떠 있는 성분으로 세지 않게 한다
    _fixg = sorted(set(d['fix_groups']))[0]
    grp = [(g if i in ok else _fixg) for i, g in enumerate(d['group'])]
    for i, q in enumerate(parts):
        if i not in ok:      # 메시를 안 담은 부품은 경계상자를 멀리 보내 후보에서 뺀다
            q['lo'] = np.full(3, 1e9)
            q['hi'] = np.full(3, 1e9 + 1.0)
    fixed_set = set()
    tb = []
    br, rej = F.bridge_floating(parts, nodes, grp, set(d['fix_groups']), fixed_set, tb,
                                d['wtol_mm'], logf=lambda m: print(m, flush=True),
                                max_gap=float(max_gap))
    print('[%s / 여유 %.1fmm / 한도 %.1fmm] 이은 다리 %d, 못 이은 쌍 %d'
          % (cand_mode, F.BRIDGE_TOL_MARGIN_MM, max_gap, len(br), len(rej)))
    for r in rej:
        if r.get('gap_mm') is not None and r['gap_mm'] <= max_gap:
            print('   못 이음:', r)
    return br, rej


if __name__ == '__main__':
    cmd = sys.argv[1] if len(sys.argv) > 1 else 'table'
    if cmd == 'req':
        make_req()
    elif cmd == 'run':
        run(label=(sys.argv[2] if len(sys.argv) > 2 else 'run'))
    elif cmd == 'table':
        table()
    elif cmd == 'replay':
        replay(*(sys.argv[2:]))
    else:
        print(__doc__)
