#!/usr/bin/env python
# -*- coding: utf-8 -*-
"""국부 재해석(2·3단계) 전용 시험대 — 전체 해석을 한 번만 풀고 캐시해 두고,
국부 재해석만 몇 번이든 다시 돌린다.

왜 필요한가: 전체 품질해석은 ccx 만 180~220초, 전체 400초다. 국부 경로를 한 줄
고칠 때마다 그걸 다시 풀면 한 번 고치는 데 7분이 든다. 전체 해석의 결과(절점·요소·
변위장·씨앗점·접합면 반력)는 국부 경로를 고쳐도 하나도 바뀌지 않으므로 한 번
풀어서 디스크에 두고 재사용한다.

  python _fea_국부시험.py global   → 전체 해석 1회 + 캐시 저장
  python _fea_국부시험.py local    → 캐시로 국부 재해석 + 수렴 판정만
  python _fea_국부시험.py repeat 3 → 같은 요청(last_request.json)을 N회 전체 해석.
      재현성(같은 설정 두 번이 다르지 않은가) 판정 전용. 브라우저를 거치지 않으므로
      '화면 클릭 순서' 변수를 제거한 상태에서 서버만의 흔들림을 재는 통제실험이다.
"""
import json
import os
import pickle
import sys
import time

import numpy as np

import fea_server as F

CACHE = os.path.join(F.WORK_ROOT, 'global_g.pkl')
OUT = os.path.join(os.path.dirname(os.path.abspath(__file__)), '_fea_국부시험_결과.json')


def _req():
    return json.load(open(os.path.join(F.WORK_ROOT, 'last_request.json'), encoding='utf-8'))


def do_global():
    q = _req()
    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'),
                  quality_bulk=True, weld_seed=True,
                  keep_arrays=True, local_refine_n=0,
                  _logf=lambda m: print(m, flush=True))
    A = r['_arrays']
    g = {'nodes': A['nodes'], 'elems': A['elems'], 'disp': A['disp'],
         'parts': A['parts'], 'seed_pairs': A['seed_pairs'],
         'iface_rf': A.get('iface_rf') or {},
         # ★ 0단계 (2026-10-12): 접합면 힘을 손으로 다시 재려면 절점력장(rfall)과
         #   '어느 절점이 어느 쌍에 묶였나'(tied_ids)가 같이 있어야 한다. 종전
         #   캐시는 합계만 들고 있어 부호·채널을 사후 검증할 수 없었다.
         'rfall': A.get('rfall'),
         'iface_tied': {k: (v.get('tied_ids') if isinstance(v, dict) else None)
                        for k, v in (A.get('iface_rf') or {}).items()},
         'h_global': float(A.get('h_global') or 10.0)}
    bk = {'material': q.get('material', F.DEF_MAT), 'E': q.get('E'), 'nu': q.get('nu'),
          'load_N': float(q.get('load_N', 1000.0)),
          'pick_radius': q.get('pick_radius'),
          'weld_mode': q.get('weld_mode') or 'bond',
          'weld_throat': float(q.get('weld_throat') or 0.0),
          'plastic': False, 'load_points': q.get('load_points'),
          'fix_points': q.get('fix_points')}
    with open(CACHE, 'wb') as f:
        pickle.dump({'g': g, 'cands': r['candidates'], 'bk': bk,
                     'mesh': r['mesh'], 'result': r['result']}, f, protocol=4)
    print('전체 해석 %.0f초 / 절점 %d / 후보 %d / 캐시 %s'
          % (time.time() - t0, r['mesh']['nodes'], len(r['candidates']), CACHE))


def _run_global(q, local_n=0, log=False):
    return 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'),
                     quality_bulk=True, weld_seed=True,
                     keep_arrays=False, local_refine_n=local_n,
                     _logf=(lambda m: print(m, flush=True)) if log else None)


REPOUT = os.path.join(os.path.dirname(os.path.abspath(__file__)),
                      '_fea_재현성시험_결과.json')


def do_repeat(n=3):
    """같은 요청을 N 회 돌려 숫자가 소수점 이하까지 같은지 본다.
    '같다/다르다' 만 적고 원인을 추측해 쓰지 않는다."""
    q = _req()
    rows = []
    for i in range(n):
        t0 = time.time()
        r = _run_global(q, local_n=0, log=False)
        m, res = r['mesh'], r['result']
        row = {'run': i + 1,
               'nodes': int(m['nodes']), 'elements': int(m['elements']),
               'welds_tied': m.get('welds'),
               'seeds': m.get('weld_seeds'),
               'parts_meshed': m.get('parts'), 'parts_failed': m.get('parts_failed'),
               'disp_mm': round(float(res['max_disp_mm']), 6),
               'struct_MPa': round(float(res['max_vonmises_MPa']), 6),
               'peak_MPa': round(float(res['peak_vonmises_MPa']), 6),
               'compliance_Nmm': r.get('compliance_Nmm'),
               'cand_at': [[round(float(x), 3) for x in c['at']]
                           for c in (r.get('candidates') or [])[:6]],
               # ★ 제외 부품 표의 원자료. 추측으로 사유를 적지 않기 위해
               #   서버가 실제로 남긴 failed_parts / auto_dropped 를 그대로 싣는다.
               'failed_parts': r.get('failed_parts'),
               'excluded_parts': r.get('excluded_parts'),
               'auto_dropped': r.get('auto_dropped'),
               'repair': (r.get('mesh') or {}).get('repair'),
               'sec': round(time.time() - t0, 1)}
        rows.append(row)
        print('[%d/%d] 절점 %s / 묶인절점 %s / 변위 %.6f / 구조최대 %.6f (%.0f초)'
              % (i + 1, n, '{:,}'.format(row['nodes']), row['welds_tied'],
                 row['disp_mm'], row['struct_MPa'], row['sec']), flush=True)
    keys = ['nodes', 'elements', 'welds_tied', 'seeds', 'parts_meshed',
            'disp_mm', 'struct_MPa', 'peak_MPa', 'compliance_Nmm', 'cand_at']
    same = {k: (len(set(json.dumps(x[k], sort_keys=True) for x in rows)) == 1)
            for k in keys}
    rep = {'n': n, 'rows': rows, 'identical': same,
           'all_identical': all(same.values())}
    json.dump(rep, open(REPOUT, 'w', encoding='utf-8'), ensure_ascii=False, indent=1)
    print('항목별 동일 여부:', json.dumps(same, ensure_ascii=False))
    print('전부 동일:', rep['all_identical'], '→', REPOUT)


def _permute(q, seed):
    """요청의 부품 순서만 바꾼다(형상·좌표·조건은 한 글자도 안 바꾼다).
    브라우저에서 '다른 부재를 클릭해 담았다' 와 서버가 보는 것이 똑같은 상태."""
    import random
    P = np.asarray(q['positions'], float).reshape(-1, 3)
    sz = [int(x) // 3 for x in q['part_sizes']]       # part_sizes = 좌표 개수(x3)
    nm = list(q['part_names'])
    off, chunks = 0, []
    for k, n in enumerate(sz):
        chunks.append((P[off:off + n], q['part_sizes'][k], nm[k]))
        off += n
    assert off == len(P), (off, len(P))
    idx = list(range(len(chunks)))
    random.Random(seed).shuffle(idx)
    q2 = dict(q)
    q2['positions'] = np.vstack([chunks[i][0] for i in idx]).reshape(-1).tolist()
    q2['part_sizes'] = [chunks[i][1] for i in idx]
    q2['part_names'] = [chunks[i][2] for i in idx]
    return q2


def do_shuffle(n=2):
    """부품 순서만 바꿔 N 회 돌린다. 숫자가 같아야 '순서 의존성이 없다' 가 증명된다.
    3회 반복(repeat)은 '서버가 결정론적이다' 만 보여 주고 순서 의존성은 못 본다."""
    q = _req()
    rows = []
    for i in range(n):
        q2 = q if i == 0 else _permute(q, 1000 + i)
        t0 = time.time()
        r = _run_global(q2, local_n=0, log=False)
        m, res = r['mesh'], r['result']
        rows.append({'run': i + 1, 'order': ('원본' if i == 0 else '섞음 seed%d' % (1000 + i)),
                     'nodes': int(m['nodes']), 'welds_tied': m.get('welds'),
                     'seeds': m.get('weld_seeds'),
                     'parts_meshed': m.get('parts'), 'parts_failed': m.get('parts_failed'),
                     'disp_mm': round(float(res['max_disp_mm']), 6),
                     'struct_MPa': round(float(res['max_vonmises_MPa']), 6),
                     'compliance_Nmm': r.get('compliance_Nmm'),
                     'cand_at': [[round(float(x), 3) for x in c['at']]
                                 for c in (r.get('candidates') or [])[:6]],
                     'sec': round(time.time() - t0, 1)})
        print('[%d/%d %s] 절점 %s / 묶인절점 %s / 변위 %.6f / 구조최대 %.6f (%.0fs)'
              % (i + 1, n, rows[-1]['order'], '{:,}'.format(rows[-1]['nodes']),
                 rows[-1]['welds_tied'], rows[-1]['disp_mm'],
                 rows[-1]['struct_MPa'], rows[-1]['sec']), flush=True)
    keys = ['nodes', 'welds_tied', 'seeds', 'parts_meshed', 'parts_failed',
            'disp_mm', 'struct_MPa', 'compliance_Nmm', 'cand_at']
    same = {k: (len(set(json.dumps(x[k], sort_keys=True) for x in rows)) == 1)
            for k in keys}
    out = REPOUT.replace('재현성시험', '순서시험')
    json.dump({'rows': rows, 'identical': same, 'all_identical': all(same.values())},
              open(out, 'w', encoding='utf-8'), ensure_ascii=False, indent=1)
    print('항목별 동일 여부:', json.dumps(same, ensure_ascii=False))
    print('전부 동일:', all(same.values()), '→', out)


def do_local(top_n=3):
    with open(CACHE, 'rb') as f:
        D = pickle.load(f)
    g, bk = D['g'], D['bk']
    t0 = time.time()
    loc = F.local_refine(D['cands'], g, bk, top_n=top_n,
                         logf=lambda m: print(m, flush=True))
    rep = {'elapsed_sec': round(time.time() - t0, 1),
           'global': {'nodes': D['mesh']['nodes'],
                      'welds_tied': D['mesh'].get('welds'),
                      'disp_mm': D['result']['max_disp_mm'],
                      'struct_MPa': D['result']['max_vonmises_MPa'],
                      'h_global_mm': g['h_global']},
           'local': loc}
    json.dump(rep, open(OUT, 'w', encoding='utf-8'), ensure_ascii=False, indent=1,
              default=lambda o: (o.tolist() if isinstance(o, np.ndarray) else str(o)))
    for rc in loc:
        print('후보%d %s — %s' % (rc['cand'], rc.get('verdict'), rc.get('why', '')))
    print('→ %s (%.0f초)' % (OUT, rep['elapsed_sec']))


if __name__ == '__main__':
    m = sys.argv[1] if len(sys.argv) > 1 else 'local'
    if m == 'global':
        do_global()
    elif m == 'shuffle':
        do_shuffle(int(sys.argv[2]) if len(sys.argv) > 2 else 2)
    elif m == 'repeat':
        do_repeat(int(sys.argv[2]) if len(sys.argv) > 2 else 3)
    else:
        do_local(int(sys.argv[2]) if len(sys.argv) > 2 else 3)
