#!/usr/bin/env python
# -*- coding: utf-8 -*-
"""1단계 '용접 씨앗점 고정' 검증 — 묶음 수가 메시 밀도와 무관한지 본다.

왜 이 시험이 따로 필요한가
  종전 방식은 메시를 만든 뒤 '표면 절점이 상대 표면에 닿았는가' 로 종속절점을
  골랐다. 그래서 메시를 촘촘하게 하면 묶인 절점이 43% 늘어났고, 그 결과가 달라진
  것이 '메시 수렴' 때문인지 '접합이 달라진' 때문인지 구분할 수 없었다
  (= 수렴 판정 불가). 씨앗점을 CAD 에서 1회만 정하면 묶음 수가 고정된다.

쓰는 법
  python _fea_씨앗점시험.py sub   → 같은 3~10부품을 절점 2배로 2단계 (묶음 수 비교)
  python _fea_씨앗점시험.py full  → 전체 146부품 품질메시 1회 (덮임 비율·후보)
"""
import json
import os
import sys
import time

import numpy as np

import fea_server as F

REQ = os.path.join(F.WORK_ROOT, 'last_request.json')
OUT = os.path.join(os.path.dirname(os.path.abspath(__file__)),
                   '_fea_씨앗점시험_결과.json')


def load():
    if not os.path.isfile(REQ):
        raise SystemExit('요청 기록이 없습니다: %s' % REQ)
    return json.load(open(REQ, encoding='utf-8'))


def subset(q, keep):
    """부품 일부만 남긴 요청을 만든다. keep = 0-base 부품 자리번호 목록."""
    P = np.asarray(q['positions'], float).reshape(-1, 3)
    ps = [int(c) // 3 for c in q['part_sizes']]
    off, segs, sizes, names = 0, [], [], []
    nm = q.get('part_names') or []
    for pi, rows in enumerate(ps):
        seg = P[off:off + rows]
        off += rows
        if pi in keep:
            segs.append(seg)
            sizes.append(rows * 3)
            names.append(nm[pi] if pi < len(nm) else str(pi + 1))
    Pn = np.vstack(segs)
    r = dict(q)
    r['positions'] = Pn.reshape(-1).tolist()
    r['part_sizes'] = sizes
    r['part_names'] = names
    return r, Pn


def one(q, **kw):
    t0 = time.time()
    r = F.run_fea(q.get('positions') or [],
                  load_N=float(q.get('load_N', 1000.0)),
                  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'), **kw)
    res, m = r['result'], r['mesh']
    sm = r.get('struct_mask') or {}
    ws = m.get('weld_seed') or {}
    nn = m['nodes']
    row = {
        'nodes': nn, 'elems': m['elements'], 'parts': m['parts'],
        'welds_tied': m.get('welds', 0),
        'seed_pairs_cad': ws.get('pairs_cad'), 'seed_pairs_mesh': ws.get('pairs_mesh'),
        'seed_pairs_both': ws.get('pairs_both'),
        'seed_total': ws.get('seeds'), 'seed_collisions': ws.get('collisions'),
        'seed_pitch_mm': ws.get('pitch_mm'),
        'welds_without_seed': ws.get('welds_without_seed'),
        'mpc_seed_nodes': sm.get('mpc_seed_nodes'),
        'mpc_excluded': sm.get('by_mpc'),
        'mpc_cover_pct': (round(100.0 * (sm.get('by_mpc') or 0) / max(nn, 1), 3)),
        'kept_pct': sm.get('kept_pct'),
        'disp': res['max_disp_mm'], 'struct_MPa': res['max_vonmises_MPa'],
        'peak_MPa': res['peak_vonmises_MPa'],
        'compliance': (r.get('compliance') or {}).get('J') if isinstance(r.get('compliance'), dict) else r.get('compliance'),
        'cand_main': len(r.get('candidates') or []),
        'cand_singular': len(r.get('candidates_singular') or []),
        'cand_pure': sum(1 for c in (r.get('candidates') or [])
                         if not c.get('singular_suspect')),
        'sec': round(time.time() - t0, 1), 'solve_sec': r['solve_sec'],
    }
    return row, r


def show(tag, row):
    print('[%s] 절점 %s / 부품 %d / 묶인절점 %s / 씨앗 %s(충돌 %s, 격자 %smm)'
          % (tag, '{:,}'.format(row['nodes']), row['parts'],
             '{:,}'.format(row['welds_tied']), row['seed_total'],
             row['seed_collisions'], row['seed_pitch_mm']))
    print('     MPC덮임 %s%% (제외절점 %s) / 변위 %.4fmm / 구조최대 %.1fMPa / %.1f초'
          % (row['mpc_cover_pct'], row['mpc_excluded'], row['disp'],
             row['struct_MPa'], row['sec']))
    print('     후보 %d(순수구조 %d) / 특이점의심목록 %d / 판정대상 %s%%'
          % (row['cand_main'], row['cand_pure'], row['cand_singular'],
             row['kept_pct']))


def main():
    mode = sys.argv[1] if len(sys.argv) > 1 else 'sub'
    q = load()
    rows = {}
    if mode == 'sub':
        # ★ 부분 묶음은 '고정점 부재 → 하중점 부재' 가 용접으로 실제로 이어져 있어야
        #   풀린다. 가까운 부품 8개를 그냥 고르면 공중에 뜬 덩어리가 나온다(실측).
        #   그래서 CAD 씨앗점 그래프에서 두 부재 사이의 최단 경로를 찾아 그 길만 쓴다.
        P = np.asarray(q['positions'], float).reshape(-1, 3)
        ps = [int(c) // 3 for c in q['part_sizes']]
        off, segs = 0, []
        for rows_n in ps:
            segs.append(P[off:off + rows_n]); off += rows_n
        cad = []
        for pi, seg in enumerate(segs):
            if len(seg) < 4:
                continue
            Vp, Fp = F.weld_mesh(seg)
            for (a, b) in F.split_components(Vp, Fp):
                cad.append({'cadV': a, 'cadF': b, 'part_no': pi + 1, 'pi': pi})
        print('CAD 덩어리 %d개 — 씨앗점 그래프를 만듭니다' % len(cad))
        seeds, _st = F.weld_seeds_cad(cad, 3.0)
        import collections
        adj = collections.defaultdict(set)
        for (i, j), v in seeds.items():
            if len(v):
                adj[i].add(j); adj[j].add(i)
        ctrs = np.array([c['cadV'].mean(axis=0) for c in cad])

        def nearest(pt):
            return int(np.argmin(np.linalg.norm(ctrs - np.asarray(pt, float), axis=1)))

        src = nearest((q.get('fix_points') or [{}])[0]['p'])
        dst = nearest((q.get('load_points') or [{}])[0]['p'])
        prev = {src: None}
        dq = collections.deque([src])
        while dq:
            x = dq.popleft()
            if x == dst:
                break
            for y in adj[x]:
                if y not in prev:
                    prev[y] = x; dq.append(y)
        if dst not in prev:
            print('고정점 부재와 하중점 부재가 용접으로 이어져 있지 않습니다 — 부분 시험 불가')
            return 1
        path = []
        x = dst
        while x is not None:
            path.append(x); x = prev[x]
        path = path[::-1]
        print('고정→하중 최단 경로 %d덩어리' % len(path))
        keep = set(cad[k]['pi'] for k in path)
        print('고른 부품(0-base): %s' % sorted(keep))
        qs, _ = subset(q, keep)
        for tag, budget in (('1단계', 120000), ('2단계', 240000)):
            try:
                row, _ = one(qs, budget_nodes=budget)
            except Exception as e:
                print('[%s] 실패: %s: %s' % (tag, type(e).__name__, e)); continue
            rows[tag] = row
            show(tag, row)
        if len(rows) == 2:
            a, b = rows['1단계'], rows['2단계']
            print('\n묶인절점 %d → %d (%s) / 씨앗 %d → %d'
                  % (a['welds_tied'], b['welds_tied'],
                     '같음 ✔' if a['welds_tied'] == b['welds_tied'] else '다름 ✘',
                     a['seed_total'], b['seed_total']))
            print('변위 %.4f → %.4f (%.2f%%) / 구조최대 %.1f → %.1f (%.2f%%)'
                  % (a['disp'], b['disp'],
                     abs(b['disp'] - a['disp']) / max(a['disp'], 1e-9) * 100,
                     a['struct_MPa'], b['struct_MPa'],
                     abs(b['struct_MPa'] - a['struct_MPa']) / max(a['struct_MPa'], 1e-9) * 100))
    else:
        row, _ = one(q, quality_bulk=True)
        rows['전체품질'] = row
        show('전체품질', row)
    json.dump(rows, open(OUT, 'w', encoding='utf-8'), ensure_ascii=False, indent=1)
    print('\n기록: %s' % OUT)
    return 0


if __name__ == '__main__':
    sys.exit(main())
