from __future__ import annotations
import json
import time
import re
from datetime import datetime
from pathlib import Path

import numpy as np
import pandas as pd
import joblib
from sklearn.model_selection import train_test_split
from sklearn.linear_model import LinearRegression
from sklearn.ensemble import RandomForestRegressor
from sklearn.metrics import mean_squared_error, r2_score

from airflow import DAG
from airflow.decorators import task
from airflow.models import Variable
from airflow.providers.amazon.aws.hooks.s3 import S3Hook


BUCKET = "sanzhar-s3"
RAW_KEY = "raw/youtube.csv"

OUT_PREFIX = Variable.get("YT_OUT_PREFIX", default_var="ml_artifacts/youtube_trends")
TRENDS_TIMEFRAME = Variable.get("YT_TRENDS_TIMEFRAME", default_var="today 1-m")
TRENDS_GEO = Variable.get("YT_TRENDS_GEO", default_var="")
KW_TOP_N = int(Variable.get("YT_KW_TOP_N", default_var="400"))
SLEEP_S = float(Variable.get("YT_TRENDS_SLEEP_S", default_var="2"))

SAMPLE_PER_BIN = int(Variable.get("YT_SAMPLE_PER_BIN", default_var="10000"))
TEST_SIZE = float(Variable.get("YT_TEST_SIZE", default_var="0.2"))
RANDOM_STATE = int(Variable.get("YT_RANDOM_STATE", default_var="42"))

BASE_DIR = Path("/opt/airflow/output/youtube_trends_ml")
BASE_DIR.mkdir(parents=True, exist_ok=True)

STOPWORDS = {
    "the","a","an","and","or","to","of","in","on","for","with","at","by",
    "official","video","trailer","full","episode","new","music","mv","vs",
    "best","from","ready","live","feat","ft","hd","4k","2024","2023",
    "season","part","game","clip","reaction"
}

def make_keyword(title: str, max_words: int = 3) -> str:
    if not isinstance(title, str):
        return ""
    t = title.lower()
    t = re.sub(r"http\\S+|www\\.\\S+", " ", t)
    t = re.sub(r"[^a-z0-9\\s]", " ", t)
    words = [w for w in t.split() if w not in STOPWORDS and len(w) > 2 and not w.isdigit()]
    return " ".join(words[:max_words])

def fetch_trends_scores(keywords, geo="", timeframe="today 1-m", sleep_s=2):
    from pytrends.request import TrendReq

    pytrends = TrendReq(hl="en-US", tz=360)
    out = {}

    for i in range(0, len(keywords), 5):
        batch = keywords[i:i+5]
        try:
            pytrends.build_payload(batch, timeframe=timeframe, geo=geo)
            it = pytrends.interest_over_time()

            if it is None or it.empty:
                for k in batch:
                    out[k] = 0.0
            else:
                for k in batch:
                    out[k] = float(it[k].mean()) if k in it.columns else 0.0

        except Exception:
            for k in batch:
                out[k] = 0.0

        time.sleep(sleep_s)

    return out

default_args = {"owner": "ec2-user", "retries": 1}

with DAG(
    dag_id="youtube_trends_ml_pipeline",
    default_args=default_args,
    start_date=datetime(2026, 2, 1),
    schedule=None,
    catchup=False,
    tags=["youtube", "ml", "s3", "pytrends"],
) as dag:

    @task
    def extract_youtube_csv_from_s3() -> str:
        local_path = str(BASE_DIR / "youTube.csv")
        s3 = S3Hook(aws_conn_id="aws_default")
        s3.get_key(RAW_KEY, bucket_name=BUCKET).download_file(local_path)
        return local_path

    @task
    def transform_and_build_ml_dataset(local_csv_path: str) -> dict:
        df = pd.read_csv(local_csv_path)

        df["publish_date"] = pd.to_datetime(df["publish_date"], dayfirst=True, errors="coerce")
        df["trending_date"] = pd.to_datetime(df["trending_date"], format="%y.%d.%m", errors="coerce")

        df["tags_count"] = df["tags"].astype(str).str.count(r"\|") + 1
        df["tags_length"] = df["tags"].astype(str).str.len()

        df["publish_hour"] = df["time_frame"].astype(str).str.split(":").str[0].astype(int)

        df = df[(df["comments_disabled"] == False) & (df["ratings_disabled"] == False)].copy()

        for col in ["ratings_disabled", "comments_disabled", "time_frame", "index", "video_error_or_removed"]:
            if col in df.columns:
                df = df.drop(columns=[col])

        df["comment_bin"] = pd.qcut(df["comment_count"], q=5, labels=False, duplicates="drop")

        balanced_df = (
            df.groupby("comment_bin", group_keys=False)
              .apply(lambda x: x.sample(n=min(SAMPLE_PER_BIN, len(x)), random_state=RANDOM_STATE))
              .copy()
        )
        balanced_df.drop(columns=["comment_bin"], inplace=True)

        balanced_df["log_comments"] = np.log1p(balanced_df["comment_count"])

        balanced_df["publish_month"] = balanced_df["publish_date"].dt.month
        balanced_df["publish_year"] = balanced_df["publish_date"].dt.year
        balanced_df["is_weekend"] = balanced_df["publish_date"].dt.weekday >= 5

        balanced_df["days_to_trend"] = (balanced_df["trending_date"] - balanced_df["publish_date"]).dt.days

        balanced_df["keyword"] = balanced_df["title"].apply(make_keyword)

        kw = (
            balanced_df["keyword"]
            .value_counts()
            .head(KW_TOP_N)
            .index
            .tolist()
        )
        kw = [k for k in kw if k]

        trend_map = fetch_trends_scores(
            kw,
            geo=TRENDS_GEO,
            timeframe=TRENDS_TIMEFRAME,
            sleep_s=SLEEP_S
        )

        balanced_df["trend_score"] = balanced_df["keyword"].map(trend_map).fillna(0.0)

        features = [
            "tags_count",
            "tags_length",
            "publish_hour",
            "publish_month",
            "days_to_trend",
            "is_weekend",
            "category_id",
            "publish_country",
            "trend_score",
        ]

        X = balanced_df[features].copy()
        y = balanced_df["log_comments"].copy()

        X = pd.get_dummies(X, columns=["publish_country"], drop_first=True)

        ml_df = X.copy()
        ml_df["log_comments"] = y.values

        ml_path = BASE_DIR / "ml_dataset.csv"
        ml_df.to_csv(ml_path, index=False)

        trend_map_path = BASE_DIR / "trend_map.json"
        trend_map_path.write_text(json.dumps(trend_map, ensure_ascii=False), encoding="utf-8")

        return {
            "ml_dataset_path": str(ml_path),
            "trend_map_path": str(trend_map_path),
            "n_keywords": len(kw),
            "timeframe": TRENDS_TIMEFRAME,
        }

    @task
    def train_models_and_metrics(info: dict) -> dict:
        ml_df = pd.read_csv(info["ml_dataset_path"])

        y = ml_df["log_comments"]
        X = ml_df.drop(columns=["log_comments"])

        X_train, X_test, y_train, y_test = train_test_split(
            X, y, test_size=TEST_SIZE, random_state=RANDOM_STATE
        )

        lr = LinearRegression()
        lr.fit(X_train, y_train)
        preds_lr = lr.predict(X_test)

        lr_rmse = float(np.sqrt(mean_squared_error(y_test, preds_lr)))
        lr_r2 = float(r2_score(y_test, preds_lr))

        rf = RandomForestRegressor(n_estimators=100, random_state=RANDOM_STATE, n_jobs=-1)
        rf.fit(X_train, y_train)
        preds_rf = rf.predict(X_test)

        rf_rmse = float(np.sqrt(mean_squared_error(y_test, preds_rf)))
        rf_r2 = float(r2_score(y_test, preds_rf))

        importances = (
            pd.Series(rf.feature_importances_, index=X.columns)
            .sort_values(ascending=False)
        )

        metrics = {
            "dataset": {
                "rows": int(ml_df.shape[0]),
                "features": int(X.shape[1]),
                "target": "log_comments",
                "timeframe": info["timeframe"],
                "n_keywords": int(info["n_keywords"]),
            },
            "models": {
                "linear_regression": {"rmse": lr_rmse, "r2": lr_r2},
                "random_forest": {"rmse": rf_rmse, "r2": rf_r2, "n_estimators": 100},
            },
            "top_feature_importances_rf": importances.head(15).to_dict(),
            "trained_at_utc": datetime.utcnow().isoformat(),
        }

        lr_path = BASE_DIR / "model_lr.joblib"
        rf_path = BASE_DIR / "model_rf.joblib"
        metrics_path = BASE_DIR / "metrics.json"

        joblib.dump(lr, lr_path)
        joblib.dump(rf, rf_path)
        metrics_path.write_text(json.dumps(metrics, indent=2, ensure_ascii=False), encoding="utf-8")

        return {
            "ml_dataset_path": info["ml_dataset_path"],
            "trend_map_path": info["trend_map_path"],
            "lr_path": str(lr_path),
            "rf_path": str(rf_path),
            "metrics_path": str(metrics_path),
        }

    @task
    def upload_all_to_s3(artifacts: dict) -> str:
        s3 = S3Hook(aws_conn_id="aws_default")

        run_id = datetime.utcnow().strftime("%Y%m%d_%H%M%S")
        prefix = f"{OUT_PREFIX}/run={run_id}"

        uploads = {
            "ml_dataset.csv": artifacts["ml_dataset_path"],
            "trend_map.json": artifacts["trend_map_path"],
            "model_lr.joblib": artifacts["lr_path"],
            "model_rf.joblib": artifacts["rf_path"],
            "metrics.json": artifacts["metrics_path"],
        }

        for name, path in uploads.items():
            s3.load_file(
                filename=path,
                key=f"{prefix}/{name}",
                bucket_name=BUCKET,
                replace=True,
            )

        return f"s3://{BUCKET}/{prefix}/"

    csv_path = extract_youtube_csv_from_s3()
    info = transform_and_build_ml_dataset(csv_path)
    artifacts = train_models_and_metrics(info)
    out = upload_all_to_s3(artifacts)