"""Reproduce the portfolio samples using only Python's standard library.
SQL and campaign inputs are synthetic. Customer records are historical public UCI data.
Source: Cardoso (2013), https://doi.org/10.24432/C5030X, CC BY 4.0.
Run from any directory: python3 path/to/verify-samples.py
"""
import csv
import hashlib
import json
import math
import sqlite3
from pathlib import Path

folder = Path(__file__).resolve().parent
with sqlite3.connect(':memory:') as db:
    db.execute('PRAGMA foreign_keys=ON')
    db.executescript((folder/'sql-quality-check.sql').read_text())
    source = db.execute('SELECT SUM(quantity*unit_price) FROM orders').fetchone()[0]
    naive = db.execute('SELECT SUM(revenue),SUM(spend),COUNT(*) FROM naive_report').fetchone()
    corrected = db.execute('SELECT SUM(revenue),SUM(spend),COUNT(*),SUM(spend_without_recorded_signups) FROM validated_report').fetchone()
    assert source == 320
    assert naive == (490, 310, 3)
    assert corrected == (320, 170, 4, 40)
    assert db.execute("SELECT revenue,spend FROM validated_report WHERE product_id='D'").fetchone() == (0,0)
    assert not db.execute('PRAGMA foreign_key_check').fetchall()

with (folder/'campaign-readout.csv').open(newline='') as stream:
    experiment = {row['variant']: row for row in csv.DictReader(stream)}
n_a, n_b = [int(experiment[key]['assigned_visitors']) for key in ['A','B']]
p_a, p_b = [int(experiment[key]['signups'])/int(experiment[key]['assigned_visitors']) for key in ['A','B']]
assert p_a == .05 and p_b == .06
diff = p_b-p_a
se = math.sqrt(p_a*(1-p_a)/n_a+p_b*(1-p_b)/n_b)
interval = [(diff-1.96*se)*100,(diff+1.96*se)*100]
assert round(diff*100,2)==1 and round(diff/p_a*100,2)==20
assert round(interval[0],2)==-.82 and round(interval[1],2)==2.82

source_path = folder/'wholesale-customers.csv'
assert hashlib.sha256(source_path.read_bytes()).hexdigest() == 'c3d018c643565b85cee733c4a2ac76dd76e080e857cb23f0ccfcc2e15a6c17ef'
with source_path.open(newline='') as stream:
    customers = [{key:int(value) for key,value in row.items()} for row in csv.DictReader(stream)]
categories = ['Fresh','Milk','Grocery','Frozen','Detergents_Paper','Delicassen']
selected = ['Milk','Grocery','Detergents_Paper']
channels = {}
for channel in [1,2]:
    group = [row for row in customers if row['Channel']==channel]
    total = sum(row[key] for row in group for key in categories)
    selected_total = sum(row[key] for row in group for key in selected)
    channels[channel] = {'count':len(group),'total':total,'selected_categories':selected_total,'share_percent':round(selected_total/total*100,1)}
assert channels[1] == {'count':298,'total':7999569,'selected_categories':2444918,'share_percent':30.6}
assert channels[2] == {'count':142,'total':6619931,'selected_categories':4871858,'share_percent':73.6}
assert sum(channel['total'] for channel in channels.values()) == 14619500
assert sum(channel['selected_categories'] for channel in channels.values()) == 7316776
print(json.dumps({'status':'passed','sql':{'source_revenue':source,'naive_report':naive,'validated_report':corrected},'synthetic_campaign':{'difference_pp':round(diff*100,2),'relative_lift_percent':round(diff/p_a*100,2),'approximate_95_interval_pp':[round(x,2) for x in interval]},'public_customer_channels':channels},indent=2))
