# Source: content/notes/ml/feature-engineering.md
# Independent CPU example; use the curriculum environment.
# See /notes/ml/#example-environment or /notes/deep-learning/#example-environment.

import numpy as np
import pandas as pd

events = pd.DataFrame({
    "user_id": ["a", "a", "a", "a", "b", "b"],
    "ts": pd.to_datetime(["2025-01-01", "2025-01-03", "2025-01-10",
                          "2025-01-11", "2025-01-02", "2025-01-06"], utc=True),
    "amount": [10.0, 20.0, 30.0, 40.0, 5.0, 15.0]})

def past_features(frame):
    result = pd.DataFrame(index=frame.index, columns=["mean7", "past_mean"], dtype=float)
    for _, group in frame.groupby("user_id", sort=False):
        ordered = group.sort_values("ts")
        assert ordered["ts"].is_unique, "Define a simultaneous-event policy first"
        series = ordered.set_index("ts")["amount"]
        result.loc[ordered.index, "mean7"] = series.rolling(
            "7D", closed="left", min_periods=1).mean().to_numpy()
        result.loc[ordered.index, "past_mean"] = series.shift(1).expanding().mean().to_numpy()
    return result

before = past_features(events)
assert np.allclose(before["mean7"], [np.nan, 10, 20, 30, np.nan, 5], equal_nan=True)
later = pd.DataFrame({"user_id": ["a"],
    "ts": pd.to_datetime(["2025-01-20"], utc=True), "amount": [999.0]})
extended = pd.concat([events, later], ignore_index=True)
after = past_features(extended).iloc[:len(events)]
assert np.allclose(before, after, equal_nan=True)

snapshots = pd.DataFrame({"user_id": ["a", "a", "b"],
    "available_at": pd.to_datetime(["2024-12-31", "2025-01-10", "2025-01-01"], utc=True),
    "known_count": [2, 9, 1]})
joined = pd.merge_asof(events.sort_values("ts"), snapshots.sort_values("available_at"),
    left_on="ts", right_on="available_at", by="user_id", direction="backward",
    allow_exact_matches=False)
assert (joined["available_at"] < joined["ts"]).all()
assert joined.loc[(joined.user_id == "a") &
                  (joined.ts == pd.Timestamp("2025-01-10", tz="UTC")), "known_count"].item() == 2
print(events.join(before).to_string(index=False))
