from metaflow import FlowSpec, step, card, conda_base, conda, current, Parameter, kubernetes
from metaflow.cards import Markdown, Table, Image
from io import BytesIO

URL = (
    "https://metaflow-demo-public.s3.us-west-2.amazonaws.com"
    "/taxi/sandbox/train_sample.parquet"
)
DAYS = ["Sunday", "Monday", "Tuesday", "Wednesday", "Thursday", "Friday", "Saturday"]

# @conda_base(libraries={"datashader": "0.14.4", "pandas": "1.5.3", "pyarrow": "8.0.0"})
@conda_base(libraries={"datashader": "0.14.0", "pandas": "1.4.2", "pyarrow": "5.0.0"})
# @conda("./test-env.yaml")
class NYCVizFlow(FlowSpec):

    data_url = Parameter("data_url", default=URL)

    @card
    @step
    def start(self):
        import pandas as pd

        df = pd.read_parquet(self.data_url)
        self.size = df.shape[0]
        print(f"Processing {self.size} datapoints...")
        df["key"] = pd.to_datetime(df["key"])
        self.dfs_by_dow = [(d, df[df.key.dt.day_name() == d]) for d in DAYS]
        self.next(self.visualize, foreach="dfs_by_dow")

    # UNCOMMENT THIS LINE FOR CLOUD EXECUTION:
    # @kubernetes
    @card
    @step
    def visualize(self):
        self.day_index = self.index
        self.dow, df = self.input
        self.dow_size = df.shape[0]

        print(f"Plotting {self.dow}")
        self.img = render_heatmap(df)
        self.next(self.validate)

    @card(type="blank")
    @step
    def validate(self, inputs):

        # validate that all data points got processed
        self.num_rows = sum(inp.dow_size for inp in inputs)
        assert inputs[0].size == self.num_rows

        # produce a report card
        current.card.append(Markdown("# NYC Taxi drop off locations by weekday"))
        rows = []
        for task in sorted(inputs, key=lambda x: x.day_index):
            rows.append([task.dow, Image(task.img)])
        current.card.append(Table(rows, headers=["Day of week", "Heatmap"]))
        self.next(self.end)

    @step
    def end(self):
        print("Success!")


def render_heatmap(df, x_range=(-74.02, -73.90), y_range=(40.69, 40.83), w=375, h=450):
    from datashader import Canvas, count, colors
    from datashader import transfer_functions as tr_fns

    cvs = Canvas(plot_width=w, plot_height=h, x_range=x_range, y_range=y_range)
    agg = cvs.points(
        df, "dropoff_longitude", "dropoff_latitude", count("passenger_count")
    )
    img = tr_fns.shade(agg, cmap=colors.Hot, how="log")
    img = tr_fns.set_background(img, "black")
    buf = BytesIO()
    img.to_pil().save(buf, format="png")
    return buf.getvalue()


if __name__ == "__main__":
    NYCVizFlow()
