use arrow::array::RecordBatch as ArrowRecordBatch;
use arrow::datatypes::Schema as ArrowSchema;
use arrow::ipc::reader::StreamReader;
use polars::error::PolarsError;
use polars::frame::DataFrame;
use polars::prelude::*;
use polars_io::ipc::{IpcCompression, IpcStreamWriter};
use std::io::Cursor;
use std::sync::Arc;

pub fn write_to_arrow(
    df: &mut DataFrame,
    compression: Option<IpcCompression>,
) -> PolarsResult<(Arc<ArrowSchema>, Vec<ArrowRecordBatch>)> {
    let mut buffer = Vec::new();
    IpcStreamWriter::new(&mut buffer)
        .with_compression(compression)
        .finish(df)?;

    let cursor = Cursor::new(buffer);
    let stream_reader = StreamReader::try_new(cursor, None)
        .map_err(|e| PolarsError::ComputeError(format!("{e}").into()))?;

    let schema = stream_reader.schema();
    let batches: Result<Vec<_>, _> = stream_reader.collect();
    let batches = batches.map_err(|e| PolarsError::ComputeError(format!("{e}").into()))?;

    Ok((schema, batches))
}

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

    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 compressed_write() {
        let mut df = create_df();
        let arrow_data = write_to_arrow(&mut df, Some(IpcCompression::ZSTD));
        let (arrow_schema, record_batches) = arrow_data.unwrap();

        let df_schema = df.schema();
        let arrow_schema_fields = arrow_schema.fields();
        let df_schema_fields: Vec<Field> = df_schema.iter_fields().collect();
        let column_names = df.get_column_names();
        let column_names_owned = df.get_column_names_owned();

        assert_eq!(arrow_schema_fields.len(), df_schema_fields.len());
        assert_eq!(arrow_schema_fields.len(), df.width());

        for (i, field) in arrow_schema_fields.iter().enumerate() {
            assert_eq!(field.name().to_owned(), column_names_owned[i].to_string());
        }

        assert!(!record_batches.is_empty());
        assert_eq!(record_batches.len(), 1);

        let batch = &record_batches[0];
        assert_eq!(batch.num_rows(), df.height());
        assert_eq!(batch.num_columns(), df.width());

        assert_eq!(column_names.len(), batch.num_columns());
        assert_eq!(column_names_owned.len(), batch.num_columns());
    }

    #[test]
    fn uncompressed_write() {
        let mut df = create_df();
        let arrow_data = write_to_arrow(&mut df, None);
        let (arrow_schema, record_batches) = arrow_data.unwrap();

        let df_schema = df.schema();
        let arrow_schema_fields = arrow_schema.fields();
        let df_schema_fields: Vec<Field> = df_schema.iter_fields().collect();
        let column_names = df.get_column_names();
        let column_names_owned = df.get_column_names_owned();

        assert_eq!(arrow_schema_fields.len(), df_schema_fields.len());
        assert_eq!(arrow_schema_fields.len(), df.width());

        for (i, field) in arrow_schema_fields.iter().enumerate() {
            assert_eq!(field.name().to_owned(), column_names_owned[i].to_string());
        }

        assert!(!record_batches.is_empty());
        assert_eq!(record_batches.len(), 1);

        let batch = &record_batches[0];
        assert_eq!(batch.num_rows(), df.height());
        assert_eq!(batch.num_columns(), df.width());

        assert_eq!(column_names.len(), batch.num_columns());
        assert_eq!(column_names_owned.len(), batch.num_columns());
    }
}