Documentation
use polars::error::PolarsResult;
use polars::frame::DataFrame;
use polars::lazy::dsl::Expr;
use polars::prelude::*;

#[derive(Clone, Debug, Default)]
pub struct FeatureOperator {
    feature_expressions: Vec<Vec<Expr>>,
}

impl FeatureOperator {
    pub fn new(feature_expressions: Vec<Vec<Expr>>) -> FeatureOperator {
        FeatureOperator {
            feature_expressions: feature_expressions,
        }
    }

    pub fn collect_lazy_expressions(
        &self,
        mut df: DataFrame,
        n_rows: Option<u32>,
        exclude_columns: Option<Vec<String>>,
    ) -> PolarsResult<DataFrame> {
        let exclude_columns = exclude_columns.unwrap_or_default();

        for expressions in &self.feature_expressions {
            df = match n_rows {
                Some(n_rows) => df
                    .lazy()
                    .with_columns(expressions)
                    .tail(n_rows)
                    .select([all().exclude(&exclude_columns)])
                    .collect()?,
                None => df
                    .lazy()
                    .with_columns(expressions)
                    .select([all().exclude(&exclude_columns)])
                    .collect()?,
            }
        }

        Ok(df)
    }

    pub fn feature_expressions_count(&self) -> usize {
        self.feature_expressions.iter().flatten().count()
    }
}

#[cfg(test)]
mod tests {
    use super::*;

    fn create_feature_operator() -> FeatureOperator {
        let expressions = vec![vec![
            (col("high") + col("low")).alias("addition"),
            (col("high") - col("low")).alias("subtraction"),
            (col("high") / col("low")).alias("division"),
            (col("high") * col("low")).alias("multiplication"),
        ]];

        FeatureOperator::new(expressions)
    }

    fn create_df() -> DataFrame {
        df!(
            "ticker" => ["AAPL", "NVDA", "MSFT", "GOOG", "AMZN"],
            "price" => [229.9, 138.93, 420.56, 166.41, 188.4],
            "high" => [231.31, 139.6, 424.04, 167.62, 189.83],
            "low" => [228.6, 136.3, 417.52, 164.78, 188.44],
        )
        .unwrap()
    }

    #[test]
    fn count_features() {
        let feature_operator = create_feature_operator();

        assert_eq!(feature_operator.feature_expressions_count(), 4);
    }

    #[test]
    fn collect_expressions() {
        let df = create_df();
        let (height, width) = df.shape();
        let feature_operator = create_feature_operator();
        let feature_count = feature_operator.feature_expressions_count();

        let df = feature_operator
            .collect_lazy_expressions(df, None, None)
            .unwrap();

        dbg!(&df);

        assert_eq!(df.shape(), (height, width + feature_count));
        assert_eq!(
            df.get_column_names(),
            &[
                "ticker",
                "price",
                "high",
                "low",
                "addition",
                "subtraction",
                "division",
                "multiplication"
            ]
        );

        let exclude_columns = vec![String::from("addition"), String::from("subtraction")];
        let df = feature_operator
            .collect_lazy_expressions(df, Some(5), Some(exclude_columns.clone()))
            .unwrap();

        assert_eq!(
            df.shape(),
            (height, width + feature_count - exclude_columns.len())
        );
        assert_eq!(
            df.get_column_names(),
            &[
                "ticker",
                "price",
                "high",
                "low",
                "division",
                "multiplication"
            ]
        );
    }

    #[test]
    fn derived_expressions() {
        let df = create_df();

        let expressions = vec![
            vec![
                (col("high") + col("low")).alias("addition"),
                (col("high") - col("low")).alias("subtraction"),
                (col("high") / col("low")).alias("division"),
                (col("high") * col("low")).alias("multiplication"),
            ],
            vec![
                (col("high").gt("low")).alias("up"),
                (col("high").lt("low")).alias("down"),
            ],
            vec![(col("up").shift(lit(1))).alias("shifted")],
        ];

        let feature_operator = FeatureOperator::new(expressions);
        let df = feature_operator
            .collect_lazy_expressions(df, None, None)
            .unwrap();

        dbg!(df);
    }

    #[test]
    fn rolling_window_expression() {
        let df = create_df();

        let rolling_window_expression = (col("price").rolling_mean(RollingOptionsFixedWindow {
            window_size: 3,
            min_periods: 3,
            weights: None,
            center: false,
            fn_params: Default::default(),
        }))
        .alias("rolling_mean");

        let expressions = vec![vec![rolling_window_expression]];
        let feature_operator = FeatureOperator::new(expressions);

        let df = feature_operator
            .collect_lazy_expressions(df, None, None)
            .unwrap();

        dbg!(&df);

        assert!(df.get_column_names_str().to_vec().contains(&"rolling_mean"));
    }
}