"""Count US executions that used pentobarbital, 2010 to 2026, from DPI yearly execution lists.

usage: python analyze.py <dir-with-YYYY.html> <output-dir>

Each https://deathpenaltyinfo.org/executions/<year> page holds one table with a
"Drug Protocol" column. Protocol labels are classified as:
  single  1-drug protocol naming pentobarbital (includes DPI's
          "1-drug (Undisclosed, likely Pentobarbital)", flagged in `likely`)
  three   3-drug protocol naming pentobarbital as the first drug
  other   every other protocol or method
"""
import csv, hashlib, html, json, re, sys
from collections import Counter, defaultdict
from pathlib import Path

src, out = Path(sys.argv[1]), Path(sys.argv[2])
YEARS = range(2010, 2027)

def clean(x):
    x = re.sub(r'<br\s*/?>', ' ', x)
    x = html.unescape(re.sub(r'<[^>]+>', '', x))
    x = x.replace('\u200b', '').replace('\u2011', '-').replace('\xa0', ' ')
    return re.sub(r'\s+', ' ', x).strip()

def classify(protocol):
    p = protocol.lower()
    if 'pentobarbital' not in p:
        return 'other'
    return 'three' if p.startswith('3') else 'single'

rows, manifest = [], []
for y in YEARS:
    raw = (src / f'{y}.html').read_bytes()
    page = raw.decode('utf-8', 'ignore')
    updated = re.search(r'Last Updated:\s*(.*?)</p>', page)
    manifest.append({
        'year': y,
        'url': f'https://deathpenaltyinfo.org/executions/{y}',
        'sha256': hashlib.sha256(raw).hexdigest(),
        'page_last_updated': clean(updated.group(1)) if updated else None,
    })
    for table in re.findall(r'<table>(.*?)</table>', page, flags=re.S):
        head = [clean(h) for h in re.findall(r'<th[^>]*>(.*?)</th>', table, flags=re.S)]
        if 'Drug Protocol' not in head:
            continue
        body = re.sub(r'<thead>.*?</thead>', '', table, flags=re.S)
        for tr in re.findall(r'<tr>(.*?)</tr>', body, flags=re.S):
            cells = [clean(c) for c in re.findall(r'<td[^>]*>(.*?)</td>', tr, flags=re.S)]
            if len(cells) != len(head):
                continue
            d = dict(zip(head, cells))
            rows.append({
                'year': y,
                'date': d['Date'],
                'jurisdiction': d['State'],
                'name': d['Name'],
                'method': d['Method'],
                'drug_protocol': d['Drug Protocol'],
                'category': classify(d['Drug Protocol']),
                'likely': 'undisclosed' in d['Drug Protocol'].lower(),
            })

out.mkdir(parents=True, exist_ok=True)
with open(out / 'executions-2010-2026.csv', 'w', newline='', encoding='utf-8') as f:
    w = csv.DictWriter(f, fieldnames=list(rows[0]))
    w.writeheader(); w.writerows(rows)

annual = defaultdict(Counter)
for r in rows:
    annual[r['year']][r['category']] += 1
with open(out / 'annual-counts.csv', 'w', newline='', encoding='utf-8') as f:
    w = csv.writer(f)
    w.writerow(['year', 'all_executions', 'single_drug_pentobarbital', 'three_drug_pentobarbital', 'any_pentobarbital', 'other'])
    for y in YEARS:
        c = annual[y]
        w.writerow([y, sum(c.values()), c['single'], c['three'], c['single'] + c['three'], c['other']])

juris = Counter(r['jurisdiction'] for r in rows if r['category'] == 'single')
with open(out / 'single-drug-by-jurisdiction.csv', 'w', newline='', encoding='utf-8') as f:
    w = csv.writer(f)
    w.writerow(['jurisdiction', 'single_drug_pentobarbital_executions'])
    for k, v in juris.most_common():
        w.writerow([k, v])

json.dump({'sources': manifest, 'rows': len(rows)}, open(out / 'sources.json', 'w'), indent=2)
print(len(rows), 'rows;', sum(juris.values()), 'single-drug;', sum(annual[y]['three'] for y in YEARS), 'three-drug')
