import math
from db import get_connection

def stats(returns):
    n = len(returns)
    if n == 0:
        return None, None, None, 0
    mean = sum(returns)/n
    if n < 2:
        return mean, 0, 0, n
    var = sum((r-mean)**2 for r in returns)/(n-1)
    std = math.sqrt(var)
    se = std/math.sqrt(n)
    return mean, std, se, n

conn = get_connection()
cursor = conn.cursor(dictionary=True)
cursor.execute('''
    SELECT t.entry_date, t.return_pct, b.final_score_v3
    FROM trade_simulation_results t
    JOIN backtest_results_v2 b
      ON b.company_id = t.company_id
     AND b.fundamentals_period = t.fundamentals_period
     AND b.snapshot_date = t.entry_date
    WHERE t.fundamentals_period = 'annual'
      AND b.final_score_v3 IS NOT NULL
''')
rows = cursor.fetchall()
cursor.close()
conn.close()

dates = [r['entry_date'] for r in rows]
min_d, max_d = min(dates), max(dates)
midpoint = min_d + (max_d - min_d) / 2
print(f'Total operaciones con score v3: {len(rows)}')
print(f'Punto medio: {midpoint}')
print()

def get_returns(score_filter=None, date_filter=None):
    out = []
    for r in rows:
        if score_filter is not None and not (float(r['final_score_v3']) > score_filter):
            continue
        if date_filter == 'first' and not (r['entry_date'] < midpoint):
            continue
        if date_filter == 'second' and not (r['entry_date'] >= midpoint):
            continue
        out.append(float(r['return_pct']))
    return out

base_first = get_returns(date_filter='first')
base_second = get_returns(date_filter='second')
mb1, sb1, seb1, nb1 = stats(base_first)
mb2, sb2, seb2, nb2 = stats(base_second)
print(f'Baseline (todas las señales) primera mitad: media={mb1:.2f}% n={nb1}')
print(f'Baseline (todas las señales) segunda mitad: media={mb2:.2f}% n={nb2}')
print()

print('=== CALIBRACION (primera mitad) con final_score_v3 ===')
best_threshold = None
best_z = -999
for threshold in [40, 45, 50, 55, 60, 65]:
    returns = get_returns(score_filter=threshold, date_filter='first')
    if len(returns) < 15:
        print(f'  >{threshold}: n={len(returns)} (insuficiente, se omite)')
        continue
    m, s, se, n = stats(returns)
    diff = m - mb1
    se_diff = math.sqrt(se**2 + seb1**2)
    z = diff/se_diff if se_diff>0 else 0
    wins = sum(1 for r in returns if r>0)
    print(f'  >{threshold}: n={n} media={m:.2f}% win_rate={wins/n*100:.1f}% z={z:.2f}')
    if z > best_z:
        best_z = z
        best_threshold = threshold
print(f'  -> Umbral elegido: >{best_threshold} (z={best_z:.2f})')
print()

print(f'=== VALIDACION fuera de muestra (segunda mitad), umbral >{best_threshold} ===')
returns2 = get_returns(score_filter=best_threshold, date_filter='second')
if len(returns2) < 5:
    print(f'  muestra insuficiente (n={len(returns2)})')
else:
    m2, s2, se2, n2 = stats(returns2)
    diff2 = m2 - mb2
    se_diff2 = math.sqrt(se2**2 + seb2**2)
    z2 = diff2/se_diff2 if se_diff2>0 else 0
    wins2 = sum(1 for r in returns2 if r>0)
    print(f'  n={n2} media={m2:.2f}% win_rate={wins2/n2*100:.1f}% z={z2:.2f}')
