"""Frozen chronological baseline/model comparison for an original synthetic case."""
from pathlib import Path
import hashlib
import json
import platform
import numpy as np
import pandas as pd
import scipy
import sklearn
from sklearn.compose import ColumnTransformer
from sklearn.dummy import DummyClassifier
from sklearn.impute import SimpleImputer
from sklearn.linear_model import LogisticRegression
from sklearn.metrics import log_loss,brier_score_loss,roc_auc_score,average_precision_score,confusion_matrix
from sklearn.pipeline import Pipeline
from sklearn.preprocessing import OneHotEncoder,StandardScaler

ROOT=Path(__file__).resolve().parent
NUMERIC=['days_since_activity','tickets_7d','tenure_days']
FEATURES=NUMERIC+['plan']
TARGET='inactive_next_7d'


def contract():return json.loads((ROOT/'experiment-contract.json').read_text(encoding='utf-8'))


def load():
    frame=pd.read_csv(ROOT/'account-snapshots.csv')
    for column in ['decision_at','features_available_at','horizon_end','label_available_at']:
        frame[column]=pd.to_datetime(frame[column],utc=True,errors='raise')
    return validate(frame)


def validate(frame):
    frame=frame.copy()
    expected=set(FEATURES+['snapshot_id','customer_id','decision_at','features_available_at',
                         'horizon_end','label_available_at','future_active_days',TARGET,'split'])
    if set(frame.columns)!=expected:raise ValueError('unexpected_schema')
    if frame.empty or frame.drop(columns=['tickets_7d']).isna().any().any():raise ValueError('missing_required_value')
    if frame['snapshot_id'].duplicated().any():raise ValueError('duplicate_snapshot')
    if frame.duplicated(['customer_id','decision_at']).any():raise ValueError('duplicate_customer_decision')
    if not frame['split'].isin(['train','validation','test']).all():raise ValueError('invalid_split')
    for column in NUMERIC+['future_active_days']:
        values=frame[column].dropna().to_numpy(dtype=float)
        if not np.isfinite(values).all() or (values<0).any() or not (values==np.floor(values)).all():
            raise ValueError('invalid_count_'+column)
    if (frame['future_active_days']>7).any():raise ValueError('active_days_exceed_horizon')
    if not (frame['features_available_at']<=frame['decision_at']).all():raise ValueError('future_feature')
    if not ((frame['horizon_end']-frame['decision_at'])==pd.Timedelta(days=7)).all():raise ValueError('wrong_horizon')
    if not ((frame['label_available_at']-frame['horizon_end'])==pd.Timedelta(days=1)).all():raise ValueError('wrong_label_delay')
    if not frame[TARGET].isin([0,1]).all():raise ValueError('invalid_target')
    if not ((frame['future_active_days']==0).astype(int)==frame[TARGET]).all():raise ValueError('target_mismatch')
    settings=contract()
    for split,freeze in [('train','train_freeze'),('validation','selection_freeze'),('test','evaluation_as_of')]:
        subset=frame[frame['split']==split]
        if subset.empty or not (subset['label_available_at']<=pd.Timestamp(settings[freeze])).all():
            raise ValueError('immature_'+split)
    for earlier,later in [('train','validation'),('validation','test')]:
        if frame.loc[frame['split']==earlier,'decision_at'].max()>=frame.loc[frame['split']==later,'decision_at'].min():
            raise ValueError('nonchronological_splits')
    return frame


def build_model(C=1.):
    numeric=Pipeline([('impute',SimpleImputer(strategy='median',add_indicator=True)),('scale',StandardScaler())])
    prepare=ColumnTransformer([('numeric',numeric,NUMERIC),
                               ('category',OneHotEncoder(handle_unknown='ignore',sparse_output=False),['plan'])],
                              remainder='drop')
    return Pipeline([('prepare',prepare),('model',LogisticRegression(C=C,max_iter=1000,solver='lbfgs'))])


def score(y,p,threshold=.5):
    y=np.asarray(y);p=np.asarray(p)
    if y.ndim!=1 or p.shape!=y.shape or not len(y) or not np.isin(y,[0,1]).all():raise ValueError('invalid_labels')
    if not np.isfinite(p).all() or ((p<0)|(p>1)).any():raise ValueError('invalid_probabilities')
    if not np.isfinite(threshold) or not 0<=threshold<=1:raise ValueError('invalid_threshold')
    tn,fp,fn,tp=confusion_matrix(y,(p>=threshold).astype(int),labels=[0,1]).ravel()
    return {'n':len(y),'positive_count':int(y.sum()),'prevalence':float(y.mean()),
            'log_loss':float(log_loss(y,p,labels=[0,1])),'brier':float(brier_score_loss(y,p)),
            'roc_auc':float(roc_auc_score(y,p)) if len(np.unique(y))==2 else None,
            'average_precision':float(average_precision_score(y,p)) if y.sum() else None,
            'threshold':threshold,'tn':int(tn),'fp':int(fp),'fn':int(fn),'tp':int(tp),
            'accuracy':float((tn+tp)/len(y)),
            'precision':float(tp/(tp+fp)) if tp+fp else None,
            'recall':float(tp/(tp+fn)) if tp+fn else None}


def run():
    frame=load();cfg=contract();sets={s:frame[frame['split']==s].copy() for s in ['train','validation','test']}
    train=sets['train'];valid=sets['validation'];test=sets['test']
    trials=[];models={}
    for C in cfg['candidate_C']:
        model=build_model(C).fit(train[FEATURES],train[TARGET]);models[C]=model
        trials.append({'C':C,'validation':score(valid[TARGET],model.predict_proba(valid[FEATURES])[:,1])})
    chosen=min(trials,key=lambda row:(row['validation']['log_loss'],row['C']))['C']
    model=models[chosen]
    baseline=DummyClassifier(strategy='prior').fit(train[FEATURES],train[TARGET])
    results={};predictions=[]
    for split in ['validation','test']:
        subset=sets[split];p=model.predict_proba(subset[FEATURES])[:,1];base=baseline.predict_proba(subset[FEATURES])[:,1]
        results[split]={'model':score(subset[TARGET],p),'baseline':score(subset[TARGET],base)}
        for (_,row),prob,bprob in zip(subset.iterrows(),p,base):
            predictions.append({'snapshot_id':row['snapshot_id'],'customer_id':row['customer_id'],'split':split,
                                'plan':row['plan'],'target':int(row[TARGET]),'probability':float(prob),
                                'baseline_probability':float(bprob)})
    report={'selected_C':chosen,'candidate_validation':trials,'results':results,
            'split_counts':{s:len(rows) for s,rows in sets.items()},
            'training_prevalence':float(train[TARGET].mean()),
            'training_numeric_medians':model.named_steps['prepare'].named_transformers_['numeric'].named_steps['impute'].statistics_.tolist(),
            'features':FEATURES,'versions':{'python':platform.python_version(),'numpy':np.__version__,
                                          'pandas':pd.__version__,'scipy':scipy.__version__,'scikit_learn':sklearn.__version__},
            'sha256':{name:hashlib.sha256((ROOT/name).read_bytes()).hexdigest() for name in
                      ['account-snapshots.csv','experiment-contract.json','make_fixture.py','evaluation_core.py']},
            'limits':cfg['scope']+' Validation selected C; test evaluated only after that fixed selection rule. No test-driven change or train+validation refit.'}
    return report,predictions,model


if __name__=='__main__':
    report,predictions,_=run()
    (ROOT/'experiment-results.json').write_text(json.dumps(report,indent=2)+'\n',encoding='utf-8')
    pd.DataFrame(predictions).to_csv(ROOT/'predictions.csv',index=False)
    print(json.dumps(report,indent=2))
