from airflow import DAG
from airflow.operators.python import PythonOperator
from airflow.operators.empty import EmptyOperator
from datetime import datetime
import sys
import os
import io
import pandas as pd
import boto3

sys.path.append('/opt/airflow')
from src.etl_logic import run_etl
from src.ml_logic import train_risk_model, train_damage_regression
from src.s3_tools import upload_df_parquet, upload_model

from src.gulia_part import extract_all_noaa, extract_aviation, transform_and_load_processed

def etl_for_map():
    df = run_etl('/opt/airflow/data/emdat.csv')
    upload_df_parquet(df, 'data/processed_risk_data.parquet')

def ml_task(model_name):
    df = run_etl('/opt/airflow/data/emdat.csv') 
    
    model_buf, cols_buf = train_risk_model(df, model_type=model_name)
    
    upload_model(model_buf, f'models/{model_name}_classifier.joblib')
    upload_model(cols_buf, f'models/{model_name}_columns.joblib')

def ml_damage_task():
    import logging
    # Читаем объединенный и очищенный датасет прямо из S3
    bucket = os.getenv('BUCKET_NAME')
    
    # ИСПРАВЛЕНИЕ: Добавлен aws_session_token
    session_token = os.getenv('AWS_SESSION_TOKEN')
    session_token = session_token.strip() if session_token else None
    
    s3 = boto3.client(
        's3',
        aws_access_key_id=os.getenv('AWS_ACCESS_KEY_ID').strip(),
        aws_secret_access_key=os.getenv('AWS_SECRET_ACCESS_KEY').strip(),
        aws_session_token=session_token,
        region_name=os.getenv('AWS_DEFAULT_REGION', 'us-east-1').strip()
    )
    
    logging.info(f"Скачиваем processed/global_risks_clean.parquet из {bucket}")
    obj = s3.get_object(Bucket=bucket, Key='processed/global_risks_clean.parquet')
    df = pd.read_parquet(io.BytesIO(obj['Body'].read()))
    
    model_buf, cols_buf = train_damage_regression(df)
    
    upload_model(model_buf, 'models/damage_regressor_classifier.joblib')
    upload_model(cols_buf, 'models/damage_regressor_columns.joblib')


with DAG('risk_analysis_v1', start_date=datetime(2023, 1, 1), schedule_interval=None, catchup=False) as dag:
    
    # --- ВЕТКА 1: Старый процесс (Обучение классификаторов) ---
    t1 = PythonOperator(task_id='prep_map', python_callable=etl_for_map)
    
    t2 = PythonOperator(task_id='train_random_forest', python_callable=ml_task, op_kwargs={'model_name': 'rf'})
    t3 = PythonOperator(task_id='train_gradient_boosting', python_callable=ml_task, op_kwargs={'model_name': 'gb'})
    t4 = PythonOperator(task_id='train_logistic_regression', python_callable=ml_task, op_kwargs={'model_name': 'lr'})
    
    t1 >> [t2, t3, t4]

    # --- ВЕТКА 2: Новый процесс (Глобальные данные ETL + Регрессия) ---
    t_extract_noaa = PythonOperator(task_id='extract_noaa_all', python_callable=extract_all_noaa)
    t_extract_avia = PythonOperator(task_id='extract_aviation', python_callable=extract_aviation)
    t_transform_load = PythonOperator(task_id='transform_and_load_to_processed', python_callable=transform_and_load_processed)
    
    # Обучение модели предсказания ущерба запускается после создания глобального датасета
    t_train_damage = PythonOperator(task_id='train_damage_regression', python_callable=ml_damage_task)
    
    [t_extract_noaa, t_extract_avia] >> t_transform_load >> t_train_damage