"""B2 非官方参考实现：描述统计、过去日期特征、日期级划分与分组检验。"""
from pathlib import Path
import json
import numpy as np
import pandas as pd
from scipy.stats import ks_2samp
from sklearn.model_selection import TimeSeriesSplit

BASE=Path(__file__).resolve().parent if '__file__' in globals() else Path.cwd()
OUT=BASE/'B2_output';OUT.mkdir(exist_ok=True)
path=next((p for p in [Path.cwd()/'B2.csv',BASE/'B2.csv',BASE.parent/'materials/B2.csv'] if p.is_file()),None)
if path is None:raise FileNotFoundError('请将 B2.csv 与参考代码放在同一目录')
df=pd.read_csv(path)
df['datetime']=pd.to_datetime(df['date']+' '+df['time'],errors='coerce')
df['pct_change']=pd.to_numeric(df['pct_change'],errors='coerce')
bad=df[['datetime','pct_change','index_id','sector','market_phase']].isna().any(axis=1) | (df['pct_change']<=-100)
quarantined=int(bad.sum());df.loc[bad].to_csv(OUT/'quarantined.csv',index=False)
df=df.loc[~bad].copy()
if df.duplicated(['index_id','datetime']).any():raise ValueError('存在同指数同时间重复记录，应确认后处理')
df=df.sort_values(['datetime','index_id']).reset_index(drop=True)
df['date_only']=df['datetime'].dt.normalize()
df['hour']=df['datetime'].dt.hour
df['time_period']=np.select([df['hour'].between(9,10),df['hour'].between(11,13)],['morning','noon'],default='afternoon')
df['up']=(df['pct_change']>0).astype(int)

# 原题要求的描述统计；不可将全量均值当作实时预测特征。
period_index_change=df.groupby(['index_id','date_only','time_period'],as_index=False).agg(avg_pct_change=('pct_change','mean'))
sector_phase_stats=df.groupby(['sector','market_phase'],as_index=False).agg(avg_pct_change=('pct_change','mean'),up_probability=('up','mean'))

# 日收益按每个时点相对上一时点的收益复合。样例按记录日期排序；
# 真实市场必须先核验交易日历、指数定义及百分比字段口径。
daily=df.groupby(['index_id','date_only'],as_index=False).agg(daily_return=('pct_change',lambda s:((1+s/100).prod()-1)*100))
daily=daily.sort_values(['index_id','date_only']).reset_index(drop=True)
daily['trend_5']=daily.groupby('index_id')['daily_return'].transform(
    lambda s:s.shift(1).rolling(5,min_periods=5).apply(lambda r:((1+r/100).prod()-1)*100,raw=True))
df=df.merge(daily[['index_id','date_only','trend_5']],on=['index_id','date_only'],how='left',validate='many_to_one')
index_sector_trend=df.groupby(['index_id','sector'],as_index=False).agg(avg_trend_5=('trend_5','mean'))

def historical_mean_by_day(frame,keys,value,new_name):
    """同一天的记录均只读取更早日期，避免同时间跨指数数据泄漏。"""
    g=frame.groupby(keys+['date_only'],as_index=False).agg(total=(value,'sum'),n=(value,'count'))
    g=g.sort_values(keys+['date_only']).reset_index(drop=True)
    previous_sum=g.groupby(keys)['total'].transform(lambda s:s.cumsum().shift(1))
    previous_count=g.groupby(keys)['n'].transform(lambda s:s.cumsum().shift(1))
    g[new_name]=previous_sum/previous_count
    return frame.merge(g[keys+['date_only',new_name]],on=keys+['date_only'],how='left',validate='many_to_one')

df=historical_mean_by_day(df,['index_id','time_period'],'pct_change','hist_period_mean')
df=historical_mean_by_day(df,['sector','market_phase'],'pct_change','hist_sector_phase_mean')

# 对日期进行 TimeSeriesSplit，同一天的所有观测一起划分。
dates=np.sort(df['date_only'].unique())
if len(dates)<6:raise ValueError('日期太少，无法使用 5 折时序划分')
splits=list(TimeSeriesSplit(n_splits=5).split(dates))
train_idx,test_idx=splits[-1]
train_data=df[df['date_only'].isin(dates[train_idx])].copy()
test_data=df[df['date_only'].isin(dates[test_idx])].copy()
assert train_data['datetime'].max()<test_data['datetime'].min()
assert set(train_data['date_only']).isdisjoint(set(test_data['date_only']))

# 此实现假设逐日更新历史特征：测试日只能读取此前已完成日的数据，
# 不更新模型参数。若预测整个未来窗口，应在预测起点冻结所有特征。
# market_phase 还需在预测时已知；若是事后标注，应滞后或用历史数据构造。
features=['hist_period_mean','hist_sector_phase_mean','trend_5']
high_quality_train=train_data.dropna(subset=features).copy()
high_quality_train.to_csv(OUT/'high_quality_train_index.csv',index=False,encoding='utf-8-sig')
test_data.to_csv(OUT/'golden_test_index.csv',index=False,encoding='utf-8-sig')
period_index_change.to_csv(OUT/'period_index_change.csv',index=False,encoding='utf-8-sig')
sector_phase_stats.to_csv(OUT/'sector_phase_stats.csv',index=False,encoding='utf-8-sig')
index_sector_trend.to_csv(OUT/'index_sector_trend.csv',index=False,encoding='utf-8-sig')

def group_ks(reference,test,keys,label):
    rows=[]
    for key,g in test.groupby(keys):
        values=key if isinstance(key,tuple) else (key,)
        ref=reference
        for field,value in zip(keys,values):ref=ref[ref[field]==value]
        row={'comparison':label,'group_type':'+'.join(keys),'group':' / '.join(map(str,values)),
             'n_reference':len(ref),'n_test':len(g)}
        if len(ref)>=2 and len(g)>=2:
            stat,pvalue=ks_2samp(ref['pct_change'],g['pct_change'],method='asymp')
            row.update(ks_stat=float(stat),p_value=float(pvalue),note='时间自相关与多重检验需额外处理')
        else:row.update(ks_stat=np.nan,p_value=np.nan,note='样本不足')
        rows.append(row)
    return pd.DataFrame(rows)

ks_tables=[]
for keys in [['sector'],['market_phase'],['sector','market_phase']]:
    ks_tables.append(group_ks(train_data,test_data,keys,'train_vs_test'))
    # 按原题提供全量对测试集的描述性比较。样本重叠，不作独立检验结论。
    overlap=group_ks(df,test_data,keys,'overall_vs_test_descriptive')
    overlap['note']='整体包含测试样本，p 值仅描述，不作独立推断'
    ks_tables.append(overlap)
ks_results=pd.concat(ks_tables,ignore_index=True)
ks_results.to_csv(OUT/'ks_results.csv',index=False,encoding='utf-8-sig')
shares=pd.DataFrame({'overall_share':df['sector'].value_counts(normalize=True),
                    'train_share':train_data['sector'].value_counts(normalize=True),
                    'test_share':test_data['sector'].value_counts(normalize=True)}).fillna(0).reset_index(names='sector')
sector_changes=int((df.groupby('index_id')['sector'].nunique()>1).sum())
weekend_dates=int(pd.Series(dates).dt.dayofweek.ge(5).sum())
summary={'valid_rows':len(df),'quarantined_rows':quarantined,'observed_dates':len(dates),
         'train_rows':len(train_data),'test_rows':len(test_data),
         'quality_train_rows':len(high_quality_train),'warmup_rows_removed':len(train_data)-len(high_quality_train),
         'train_last_date':str(train_data['date_only'].max().date()),'test_first_date':str(test_data['date_only'].min().date()),
         'index_ids_with_multiple_sectors':sector_changes,'weekend_dates_in_sample':weekend_dates,
         'ks_comparisons':len(ks_results)}
(OUT/'summary.json').write_text(json.dumps(summary,ensure_ascii=False,indent=2),encoding='utf-8')
print('运行摘要\n'+json.dumps(summary,ensure_ascii=False,indent=2))
print('\n时段特征（前 6 行）\n'+period_index_change.head(6).to_string(index=False))
print('\n板块×阶段均值及上涨概率\n'+sector_phase_stats.to_string(index=False))
print('\n板块占比\n'+shares.to_string(index=False))
print('\n训练／测试 KS 分组结果\n'+ks_results[ks_results['comparison']=='train_vs_test'].to_string(index=False))

def tbl(frame):return '<div class="table">'+frame.to_html(index=False,border=0,float_format=lambda x:f'{x:.4f}')+'</div>'
body='<h1>B2 参考代码运行结果</h1><p>非官方参考实现；由项目附件实际运行。未训练预测模型，不含模型准确率。</p>'
body+='<h2>划分与质量</h2>'+tbl(pd.DataFrame(list(summary.items()),columns=['项目','实测值']))
body+='<p>样例日期含周末，同一指数在样例中有多个板块标签。此处保留数据并按观测日期演示，真实市场需先核验交易日历及板块标签。</p>'
body+='<h2>指数×日期×时段均值（示例）</h2>'+tbl(period_index_change.head(12))
body+='<h2>板块×市场阶段均值</h2>'+tbl(sector_phase_stats)
body+='<h2>指数×板块历史五日趋势（示例）</h2>'+tbl(index_sector_trend.head(12))
body+='<h2>板块占比</h2>'+tbl(shares)
body+='<h2>训练与测试分组 KS</h2>'+tbl(ks_results[ks_results['comparison']=='train_vs_test'])
body+='<p>p≥0.05 不证明同分布。时序观测并非完全独立，多组检验也有多重比较问题；应结合效应量、样本量和按日期分块验证。</p>'
body+='<h2>整体与测试描述性 KS（原题要求）</h2>'+tbl(ks_results[ks_results['comparison']=='overall_vs_test_descriptive'])
body+='<p>整体包含测试样本，因此上述 p 值不能作为独立两样本推断。不要据此宣称测试集绝对代表整体，也不要通过混入未来样本来提高 p 值。</p>'
body+='<h2>特征使用边界</h2><p>全量描述均值单独导出。预测特征只包含更早日期信息，五日趋势排除当日。前期历史不足的训练行被排除；逐日评估时允许读入此前已发生的数据，固定多日预测则必须冻结预测起点可见数据。</p>'
body+='<p>market_phase 必须在预测时可观测。事后划定的牛熊阶段不能直接作为预测输入，应滞后或用历史规则构造；选择实际模型输入时要用明确的特征白名单，不能把原始结果列或未来标签全部送入模型。</p>'
style='body{font:17px/1.8 system-ui,"Microsoft YaHei",sans-serif;color:#17263b;max-width:1000px;margin:auto;padding:20px}table{border-collapse:collapse;width:100%}td,th{border-bottom:1px solid #dce4ee;padding:10px;text-align:left}.table{overflow:auto}h1{font-size:1.6rem}h2{font-size:1.25rem}'
(BASE/'B2_reference.html').write_text('<!doctype html><html lang="zh-CN"><meta charset="utf-8"><meta name="viewport" content="width=device-width,initial-scale=1"><title>B2 参考运行结果</title><style>'+style+'</style><body><a href="../index.html#b2">返回学习资料</a>'+body+'</body></html>',encoding='utf-8')
