# /// script
# requires-python = ">=3.13,<3.14"
# dependencies = ["polars==1.34.0"]
# ///

import argparse
import datetime as dt
import hashlib
import json
from pathlib import Path

import polars as pl

parser = argparse.ArgumentParser()
parser.add_argument('--input', type=Path, required=True)
parser.add_argument('--output', type=Path, required=True)
args = parser.parse_args()
expected_hash = 'def12dc84a8ce6cbaee42d9cc31017c171a7f2a72a43f54efa11ddfb37ce4b62'
assert hashlib.sha256(args.input.read_bytes()).hexdigest() == expected_hash, 'Input differs from the published snapshot'
raw = pl.read_csv(args.input, schema_overrides={'time': pl.String, 'id': pl.String})
required = ['time', 'id', 'mag', 'type', 'net']
assert raw.select(pl.any_horizontal(pl.col(required).is_null()).any()).item() is False
assert raw['id'].n_unique() == raw.height
assert raw['type'].unique().to_list() == ['earthquake']
assert raw['net'].unique().to_list() == ['us']
assert raw.select((pl.col('mag') >= 5).all()).item()
events = raw.with_columns(pl.col('time').str.to_datetime(time_zone='UTC').alias('origin'))
start = dt.datetime(2011, 3, 1, tzinfo=dt.UTC)
end = dt.datetime(2011, 4, 1, tzinfo=dt.UTC)
events = events.filter((pl.col('origin') >= start) & (pl.col('origin') < end)).sort('origin')
days = pl.DataFrame({'date': pl.date_range(start.date(), end.date(), interval='1d', closed='left', eager=True)})
daily = days.join(events.group_by(pl.col('origin').dt.date().alias('date')).agg(pl.len().alias('count')), on='date', how='left').with_columns(pl.col('count').fill_null(0)).sort('date')
assert daily.height == 31 and daily['count'].sum() == events.height
gaps = events.select(pl.col('origin').diff().dt.total_milliseconds().truediv(60000).alias('minutes')).slice(1)
assert gaps.height == events.height - 1 and gaps.select((pl.col('minutes') >= 0).all()).item()
mean = daily['count'].mean()
variance = daily['count'].var(ddof=1)
fitted = daily.with_columns(pl.when(pl.col('date') < dt.date(2011, 3, 11)).then(pl.lit('before')).otherwise(pl.lit('after')).alias('segment'))
fitted = fitted.with_columns(pl.col('count').mean().over('segment').alias('fitted'))
def deviance(frame, expected):
    return frame.select((2 * (pl.when(pl.col('count') > 0).then(pl.col('count') * (pl.col('count') / expected).log()).otherwise(0) - pl.col('count') + expected)).sum()).item()

result = {
    'sourceSha256': expected_hash,
    'inputRows': raw.height,
    'excludedAtBoundary': raw.height - events.height,
    'eventCount': events.height,
    'daily': json.loads(daily.with_columns(pl.col('date').dt.to_string('%Y-%m-%d')).write_json()),
    'gapsMinutes': gaps['minutes'].to_list(),
    'summary': {
        'mean': mean, 'sampleVariance': variance, 'dispersionRatio': variance / mean,
        'constantDeviance': deviance(daily, pl.lit(mean)),
        'splitDeviance': deviance(fitted, pl.col('fitted')),
        'beforeRate': fitted.filter(pl.col('segment') == 'before')['count'].mean(),
        'afterRate': fitted.filter(pl.col('segment') == 'after')['count'].mean(),
        'firstBoundaryGapMinutes': (events['origin'][0] - start).total_seconds() / 60,
        'lastBoundaryGapMinutes': (end - events['origin'][-1]).total_seconds() / 60,
        'tiedOriginGaps': gaps.filter(pl.col('minutes') == 0).height,
    },
    'missingByColumn': raw.null_count().to_dicts()[0],
    'magnitudeTypes': raw['magType'].unique().sort().to_list(),
    'magnitudeSources': raw['magSource'].unique().sort().to_list(),
}
args.output.parent.mkdir(parents=True, exist_ok=True)
args.output.write_text(json.dumps(result, indent=2) + chr(10))
print(json.dumps({'eventCount': result['eventCount'], 'summary': result['summary']}))
