import sys
import time
import pyarrow as pa
from datafusion import SessionContext, col
import datafusion.functions as F


def main(input_dir, output):
    weather_schema = pa.schema(
        [
            ("station", pa.string()),
            ("date", pa.string()),
            ("observation", pa.string()),
            ("value", pa.int32()),
            ("mflag", pa.string()),
            ("qflag", pa.string()),
            ("sflag", pa.string()),
            ("obstime", pa.string()),
        ]
    )
    ctx = SessionContext()
    weather = ctx.read_csv(input_dir, has_header=False, schema=weather_schema)
    good_rows = weather.filter(
        col("qflag").is_null()
        & F.starts_with(col("station"), "CA")
        & (col("observation") == "TMAX")
    )
    filtered_weather = good_rows.select(
        col("station"),
        F.to_date(col("date"), "%Y%m%d").alias("date"),
        (col("value") / 10.0).alias("tmax"),
    )
    maxtemp = filtered_weather.aggregate(
        col("date"), F.max(col("tmax")).alias("max_temp")
    ).sort(col("date"))
    maxtemp.write_json(output)


if __name__ == "__main__":
    inputs = sys.argv[1]
    output = sys.argv[2]
    st = time.time()
    main(inputs, output)
    en = time.time()
    print(en - st)
