Documentation
use polars::datatypes::AnyValue;
use polars::datatypes::PlSmallStr;
use polars::error::PolarsResult;
use polars::frame::DataFrame;
use polars::frame::column::Column;
use polars::prelude::Schema;
use polars::prelude::datatypes::Field;
use polars::prelude::*;
use std::collections::HashMap;

use crate::feature_operator::FeatureOperator;

#[derive(Clone, Debug, Default)]
pub struct DataFrameSchematic {
    pub schema: Schema,
    pub feature_operator: FeatureOperator,
    pub minimum_rows: Option<u32>,
}

impl DataFrameSchematic {
    pub fn new(field_types: Vec<Field>, expressions: Vec<Vec<Expr>>) -> DataFrameSchematic {
        DataFrameSchematic {
            schema: Schema::from_iter(field_types),
            feature_operator: FeatureOperator::new(expressions),
            minimum_rows: None,
        }
    }

    pub fn with_minimum_rows(&mut self, minimum_rows: u32) -> &DataFrameSchematic {
        self.minimum_rows = Some(minimum_rows);

        self
    }

    pub fn prepare_df(&self, data_point: &Vec<(PlSmallStr, AnyValue)>) -> PolarsResult<DataFrame> {
        let mut df = DataFrame::empty_with_schema(&self.schema);

        let columns: Vec<Column> = data_point
            .iter()
            .map(|(col_name, value)| Series::new(col_name.clone(), [value.clone()]).into())
            .collect();

        let new_df = DataFrame::new(columns)?;
        df.vstack_mut(&new_df)?;

        Ok(df)
    }

    pub fn prepare_data(
        &self,
        data_point: &Vec<(PlSmallStr, AnyValue)>,
    ) -> PolarsResult<DataFrame> {
        let column_map: HashMap<PlSmallStr, Series> = data_point
            .iter()
            .map(|(col_name, value)| {
                {
                    (
                        col_name.to_owned(),
                        Series::new(col_name.to_owned(), [value.clone()]),
                    )
                }
            })
            .collect();

        self.prepare_df_from_series(&column_map, 1)
    }

    pub fn prepare_df_from_columns(
        &self,
        column_data_map: &HashMap<PlSmallStr, Column>,
        series_data_length: usize,
    ) -> PolarsResult<DataFrame> {
        let mut columns = Vec::with_capacity(self.schema.len());
        for series_name in self.schema.iter_names() {
            if let Some(column) = column_data_map.get(series_name) {
                columns.push(column.to_owned());
            } else {
                let null_series =
                    Series::new_null(PlSmallStr::from_str(series_name), series_data_length);
                columns.push(null_series.into_column());
            }
        }

        DataFrame::new(columns)
    }

    pub fn prepare_df_from_series(
        &self,
        series_data_map: &HashMap<PlSmallStr, Series>,
        series_data_length: usize,
    ) -> PolarsResult<DataFrame> {
        let mut columns = Vec::with_capacity(self.schema.len());
        for series_name in self.schema.iter_names() {
            if let Some(series) = series_data_map.get(series_name) {
                columns.push(series.to_owned().into_column());
            } else {
                let null_series =
                    Series::new_null(PlSmallStr::from_str(series_name), series_data_length);
                columns.push(null_series.into_column());
            }
        }

        DataFrame::new(columns)
    }
}

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

    fn create_df_schematic() -> DataFrameSchematic {
        let field_types = vec![
            Field::new("a".into(), DataType::Float64),
            Field::new("b".into(), DataType::Float64),
            Field::new("c".into(), DataType::Float64),
        ];
        let expressions: Vec<Vec<Expr>> = Vec::new();
        DataFrameSchematic::new(field_types, expressions)
    }

    #[test]
    fn schema_creation() {
        let schema = Schema::from_iter(vec![
            Field::new("a".into(), DataType::Float64),
            Field::new("b".into(), DataType::Float64),
            Field::new("c".into(), DataType::Float64),
        ]);

        let field_types = vec![
            Field::new("a".into(), DataType::Float64),
            Field::new("b".into(), DataType::Float64),
            Field::new("c".into(), DataType::Float64),
        ];
        let expressions: Vec<Vec<Expr>> = Vec::new();
        let df_schematic = DataFrameSchematic::new(field_types, expressions);

        let mut schema_one_feature_diff = Vec::new();
        let mut schema_two_feature_diff = Vec::new();
        schema.field_compare(
            &df_schematic.schema,
            &mut schema_one_feature_diff,
            &mut schema_two_feature_diff,
        );

        assert_eq!(schema_one_feature_diff, schema_two_feature_diff);
        assert!(schema_one_feature_diff.len() == 0);
        assert!(schema_two_feature_diff.len() == 0);
    }

    #[test]
    fn prepare_and_append_df() {
        let df_schematic = create_df_schematic();
        let mut df = DataFrame::empty_with_schema(&df_schematic.schema);

        let a_field_data = Series::new("a".into(), &[1.0, 2.0, 3.0]);
        let b_field_data = Series::new("b".into(), &[10.0, 20.0, 30.0]);
        let c_field_data = Series::new_null("c".into(), 3);

        dbg!(&df_schematic.schema);

        let new_df = DataFrame::new(vec![
            a_field_data.into(),
            b_field_data.into(),
            c_field_data.into(),
        ])
        .unwrap();

        df.vstack_mut(&new_df).unwrap();

        dbg!(&df);

        let (initial_height, initial_width) = df.shape();

        let data_points: Vec<(PlSmallStr, AnyValue)> = vec![
            ("a".into(), AnyValue::Float64(4.0)),
            ("b".into(), AnyValue::Float64(40.0)),
            ("c".into(), AnyValue::Float64(400.0)),
        ];
        let new_df = df_schematic.prepare_df(&data_points).unwrap();

        df.vstack_mut(&new_df).unwrap();

        let (new_height, new_width) = df.shape();

        dbg!(&df);

        assert!(new_height > initial_height);
        assert_eq!(new_width, initial_width);
    }

    #[test]
    fn prepare_and_append_data_points() {
        let df_schematic = create_df_schematic();
        let mut df = DataFrame::empty_with_schema(&df_schematic.schema);

        let a_field_data = Series::new("a".into(), &[1.0, 2.0, 3.0]);
        let b_field_data = Series::new("b".into(), &[10.0, 20.0, 30.0]);
        let c_field_data = Series::new_null("c".into(), 3);

        let new_df = DataFrame::new(vec![
            a_field_data.into(),
            b_field_data.into(),
            c_field_data.into(),
        ])
        .unwrap();

        df.vstack_mut(&new_df).unwrap();

        let (initial_height, initial_width) = df.shape();

        dbg!("before append", &df);

        let data_points: Vec<(PlSmallStr, AnyValue)> = vec![
            ("a".into(), AnyValue::Float64(4.0)),
            ("b".into(), AnyValue::Float64(40.0)),
            ("c".into(), AnyValue::Float64(400.0)),
        ];
        let new_df = df_schematic.prepare_data(&data_points).unwrap();

        df.vstack_mut(&new_df).unwrap();

        let (new_height, new_width) = df.shape();

        dbg!("after append", &df);

        assert!(new_height > initial_height);
        assert_eq!(new_width, initial_width);
    }

    #[test]
    fn prepare_and_append_df_from_series() {
        let df_schematic = create_df_schematic();

        let data_points: Vec<(PlSmallStr, Vec<AnyValue>)> = vec![
            (
                PlSmallStr::from_str("a"),
                vec![
                    AnyValue::Float64(4.0),
                    AnyValue::Float64(4.0),
                    AnyValue::Float64(4.0),
                ],
            ),
            (
                PlSmallStr::from_str("b"),
                vec![
                    AnyValue::Float64(4.0),
                    AnyValue::Float64(4.0),
                    AnyValue::Float64(4.0),
                ],
            ),
            (
                PlSmallStr::from_str("c"),
                vec![
                    AnyValue::Float64(4.0),
                    AnyValue::Float64(4.0),
                    AnyValue::Float64(4.0),
                ],
            ),
        ];

        let series_map: HashMap<PlSmallStr, Series> = data_points
            .iter()
            .map(|(col_name, data)| {
                {
                    (
                        col_name.clone(),
                        Series::new(col_name.clone(), data.clone()),
                    )
                }
                .into()
            })
            .collect();

        let prepared_df = df_schematic
            .prepare_df_from_series(&series_map, data_points.len())
            .unwrap();

        assert!(prepared_df.shape() == (3, 3));
    }

    #[test]
    fn prepare_and_append_df_from_column() {
        let df_schematic = create_df_schematic();

        let data_points: Vec<(PlSmallStr, Vec<AnyValue>)> = vec![
            (
                PlSmallStr::from_str("a"),
                vec![
                    AnyValue::Float64(4.0),
                    AnyValue::Float64(4.0),
                    AnyValue::Float64(4.0),
                ],
            ),
            (
                PlSmallStr::from_str("b"),
                vec![
                    AnyValue::Float64(4.0),
                    AnyValue::Float64(4.0),
                    AnyValue::Float64(4.0),
                ],
            ),
            (
                PlSmallStr::from_str("c"),
                vec![
                    AnyValue::Float64(4.0),
                    AnyValue::Float64(4.0),
                    AnyValue::Float64(4.0),
                ],
            ),
        ];

        let column_map: HashMap<PlSmallStr, Column> = data_points
            .iter()
            .map(|(col_name, data)| {
                {
                    (
                        col_name.clone(),
                        Column::new(col_name.clone(), data.clone()),
                    )
                }
                .into()
            })
            .collect();

        let prepared_df = df_schematic
            .prepare_df_from_columns(&column_map, data_points.len())
            .unwrap();

        assert!(prepared_df.shape() == (3, 3));
    }
}