import json
from pandas import DataFrame
from typing import List

from app.schemas.formula import Formula, ValueType, ComputedColumn
from app.core.formula.pandas.helpers import apply_computed_columns
from app.workers.md.md_nodes.md_node import MDNode, ModelResult
from app.workers.md.md_nodes.node_utils import copy_variables
from app import schemas, models

from sqlalchemy.orm import Session
from sqlalchemy import and_

class Transform(MDNode):
  
   def run(
       self,
       db: Session,
       node: models.md.TNode,
       df: DataFrame,
       vars_in: List[schemas.md.MDVariable],
       vars_out: List[schemas.md.MDVariable],
       params
   ):
       transform_param = params['TRANSFORM_TABLE']['value']
       computed_columns = [
           ComputedColumn(
               guid=column['ID'],
               name=column['Attribute'],
               logical_type="Numeric" if column['Level']['id'] == 'Interval' \
                   else "String",
               formula=Formula.parse_obj(column['Formula'])
           )
           for column in transform_param['rows']
       ]
       nom_vars = {
           var.var_name: ValueType.String
           for var  in filter(
               lambda v: v.var_level in ["Nominal", "Binary"],
               vars_in)
       }
       int_vars = {
           var.var_name: ValueType.Number
           for var  in filter(
               lambda v: v.var_level == "Interval",
               vars_in)
       }
       transformed_df = apply_computed_columns(
           columns={**nom_vars,**int_vars},
           df=df,
           computed_columns=computed_columns
       )
       sample_table = self.create_result_table_from_dataframe(
           df=transformed_df,
           table_name='Пример данных',
           filter_rows=100
       )
       transform_stats = {
           "tables": [sample_table]
       }
       return transformed_df, ModelResult(
           df = transformed_df,
           stats = transform_stats,
           mm_code = "",
           pickle_bytes = None
       )

   def check_status(
       self,
       vars_out
   ):
      pass

   def update_metadata(
       self,
       vars_out,
       diff_out,
       diff_in,
       node_out=None
   ):
       copy_variables (self.db, self.node_orm, vars_out)

def create_vars_for_transform( db: Session,
   transform_table_param: models.md.TParameterValue,
   node_orm: models.md.TNode,
):
   vars_in = db.query(models.md.TVariable).filter(
       and_(
           models.md.TVariable.node_guid == transform_table_param.node_guid,
           models.md.TVariable.var_type == 'IN'
       )
   ).all()
   # recreating basic variables of node
   copy_variables(db, node_orm, vars_in)
   value_orm: models.md.TParameterValue = db.query(models.md.TParameterValue).filter(
       and_(
           models.md.TParameterValue.node_guid == node_orm.node_guid,
           models.md.TParameterValue.parameter_id == 'TRANSFORM_TABLE'
       )
   ).first()
   vars_dict = json.loads(value_orm.value_json)
   for var_out in vars_dict['value']['rows']:
       transform_var_out: models.md.TVariable = models.md.TVariable(
           node_guid=node_orm.node_guid,
           pipeline_guid=node_orm.pipeline_guid,
           var_type='OUT',
           var_guid=var_out['ID'],
           var_name=var_out['Attribute'],
           var_level=var_out['Level']['id'],
           var_role=var_out['Roles']['id'],
       )
       db.add(transform_var_out)
   db.commit()