from __future__ import annotations

from collections import defaultdict
from datetime import datetime, time, timezone
from statistics import median
from typing import Any

from sqlmodel import Session, func, select

from app.models.core import (
    AnalyticsDailyAggregate,
    AnalyticsImportBatch,
    AnalyticsImportFile,
    AnalyticsPostSnapshot,
    ExternalPost,
    PostDraft,
    Provider,
    Theme,
    VideoAsset,
)
from app.services.analytics_youtube_csv import (
    UNKNOWN,
    YOUTUBE_CHART,
    YOUTUBE_DAILY_TOTALS,
    YOUTUBE_VIDEO_TABLE,
    compute_snapshot_metrics,
    file_hash,
    parse_csv_by_type,
)


def as_datetime(value):
    if value is None:
        return None
    if isinstance(value, datetime):
        return value
    return datetime.combine(value, time.min, tzinfo=timezone.utc)


def safe_delta(new, old):
    if new is None or old is None:
        return None
    return new - old


def import_youtube_csv_files(
    session: Session,
    files: list[tuple[str, bytes]],
    provider: str = "youtube",
    account_id: str | None = None,
    snapshot_at: datetime | None = None,
    source_label: str = "manual",
    is_demo: bool = False,
    import_batch_label: str | None = None,
) -> dict[str, Any]:
    provider_enum = Provider(provider)
    snapshot_at = snapshot_at or datetime.now(timezone.utc)
    batch = AnalyticsImportBatch(
        provider=provider_enum,
        account_id=account_id,
        snapshot_at=snapshot_at,
        status="importing",
        total_files=len(files),
        source_label=source_label,
        is_demo=is_demo,
        import_batch_label=import_batch_label or source_label,
    )
    session.add(batch)
    session.commit()
    session.refresh(batch)

    warnings: list[str] = []
    errors: list[str] = []
    total_rows = 0
    created_posts = 0
    snapshots = 0
    daily_metrics = 0
    date_starts = []
    date_ends = []

    for filename, content in files:
        digest = file_hash(content)
        parsed = parse_csv_by_type(content)
        total_rows += parsed.row_count
        if parsed.date_range_start:
            date_starts.append(parsed.date_range_start)
        if parsed.date_range_end:
            date_ends.append(parsed.date_range_end)
        duplicate = session.exec(
            select(AnalyticsImportFile).where(
                AnalyticsImportFile.file_hash == digest,
                AnalyticsImportFile.detected_type == parsed.detected_type,
                AnalyticsImportFile.status == "imported",
            )
        ).first()
        import_file = AnalyticsImportFile(
            batch_id=batch.id,
            file_name=filename,
            file_hash=digest,
            detected_type=parsed.detected_type,
            row_count=parsed.row_count,
            warnings_json=list(parsed.warnings),
            errors_json=list(parsed.errors),
            status="skipped_duplicate" if duplicate else "importing",
        )
        session.add(import_file)
        session.commit()
        session.refresh(import_file)
        if duplicate:
            import_file.skipped_rows = parsed.row_count
            warnings.append(f"{filename}: duplicate file hash skipped")
            import_file.warnings_json = [*import_file.warnings_json, "Duplicate file hash skipped"]
            session.add(import_file)
            session.commit()
            continue
        if parsed.detected_type == UNKNOWN:
            import_file.status = "failed"
            import_file.skipped_rows = parsed.row_count
            errors.extend([f"{filename}: {e}" for e in parsed.errors])
            session.add(import_file)
            session.commit()
            continue
        if parsed.detected_type == YOUTUBE_DAILY_TOTALS:
            imported = 0
            for row in parsed.rows:
                metric_date = as_datetime(row.get("metric_date"))
                if metric_date is None:
                    continue
                session.add(AnalyticsDailyAggregate(
                    batch_id=batch.id,
                    provider=provider_enum,
                    account_id=account_id,
                    metric_date=metric_date,
                    views=row.get("views"),
                    raw_json=row.get("raw") or {},
                ))
                imported += 1
            daily_metrics += imported
            import_file.imported_rows = imported
            import_file.skipped_rows = parsed.row_count - imported
            import_file.status = "imported"
        elif parsed.detected_type == YOUTUBE_VIDEO_TABLE:
            if parsed.total_row:
                warnings.append(
                    f"{filename}: total row views={parsed.total_row.get('views')} watch_time={parsed.total_row.get('watch_time_hours')} subscribers={parsed.total_row.get('subscribers_delta')} impressions={parsed.total_row.get('impressions')} ctr={parsed.total_row.get('impression_ctr_pct')}"
                )
            imported = 0
            for row in parsed.rows:
                external_id = row.get("external_post_id") or f"title:{row.get('title','untitled')}"
                post = session.exec(
                    select(ExternalPost).where(
                        ExternalPost.provider == provider_enum,
                        ExternalPost.external_post_id == external_id,
                        ExternalPost.account_id == account_id,
                    )
                ).first()
                if not post:
                    post = ExternalPost(
                        provider=provider_enum,
                        external_post_id=external_id,
                        account_id=account_id,
                        title=row.get("title"),
                        published_at=as_datetime(row.get("published_at")),
                        duration_sec=row.get("duration_sec"),
                        source_label=source_label,
                        is_demo=is_demo,
                        import_batch_label=import_batch_label or source_label,
                        source_file_name=filename,
                    )
                    linked = match_video_asset(session, row.get("title"), external_id)
                    if linked:
                        post.video_asset_id = linked.id
                        post.topic_id = linked.topic_id
                        post.series = linked.series
                    session.add(post)
                    session.commit()
                    session.refresh(post)
                    created_posts += 1
                else:
                    post.title = row.get("title") or post.title
                    post.published_at = as_datetime(row.get("published_at")) or post.published_at
                    post.duration_sec = row.get("duration_sec") or post.duration_sec
                    post.source_label = source_label
                    post.is_demo = is_demo
                    post.import_batch_label = import_batch_label or source_label
                    post.source_file_name = filename
                    post.updated_at = datetime.now(timezone.utc)
                    session.add(post)
                    session.commit()
                previous = session.exec(
                    select(AnalyticsPostSnapshot)
                    .where(
                        AnalyticsPostSnapshot.provider == provider_enum,
                        AnalyticsPostSnapshot.external_post_id == external_id,
                        AnalyticsPostSnapshot.account_id == account_id,
                    )
                    .order_by(AnalyticsPostSnapshot.snapshot_at.desc())
                ).first()
                derived = compute_snapshot_metrics(row, snapshot_at)
                snapshot = AnalyticsPostSnapshot(
                    batch_id=batch.id,
                    provider=provider_enum,
                    account_id=account_id,
                    external_post_id=external_id,
                    video_asset_id=post.video_asset_id,
                    snapshot_at=snapshot_at,
                    title=row.get("title"),
                    published_at=as_datetime(row.get("published_at")),
                    duration_sec=row.get("duration_sec"),
                    views=row.get("views"),
                    watch_time_hours=row.get("watch_time_hours"),
                    subscribers_delta=row.get("subscribers_delta"),
                    impressions=row.get("impressions"),
                    impression_ctr_pct=row.get("impression_ctr_pct"),
                    avg_view_duration_sec=derived.get("avg_view_duration_sec"),
                    retention_proxy_pct=derived.get("retention_proxy_pct"),
                    views_per_impression=derived.get("views_per_impression"),
                    subscribers_per_1000_views=derived.get("subscribers_per_1000_views"),
                    views_per_day_since_publish=derived.get("views_per_day_since_publish"),
                    raw_json=row.get("raw") or {},
                    source_label=source_label,
                    is_demo=is_demo,
                    import_batch_label=import_batch_label or source_label,
                    source_file_name=filename,
                )
                apply_deltas(snapshot, previous)
                session.add(snapshot)
                imported += 1
                snapshots += 1
            import_file.imported_rows = imported
            import_file.skipped_rows = parsed.row_count - imported - (1 if parsed.total_row else 0)
            import_file.status = "imported"
        elif parsed.detected_type == YOUTUBE_CHART:
            import_file.imported_rows = len(parsed.rows)
            import_file.skipped_rows = max(0, parsed.row_count - len(parsed.rows))
            import_file.status = "empty_valid" if parsed.row_count == 0 else "imported"
        import_file.warnings_json = [*import_file.warnings_json, *parsed.warnings]
        import_file.errors_json = [*import_file.errors_json, *parsed.errors]
        warnings.extend([f"{filename}: {w}" for w in parsed.warnings])
        errors.extend([f"{filename}: {e}" for e in parsed.errors])
        session.add(import_file)
        session.commit()

    if date_starts:
        batch.date_range_start = as_datetime(min(date_starts))
    if date_ends:
        batch.date_range_end = as_datetime(max(date_ends))
    batch.status = "failed" if errors and not (snapshots or daily_metrics) else "imported_with_warnings" if warnings or errors else "imported"
    batch.total_rows = total_rows
    batch.warnings_json = warnings
    batch.errors_json = errors
    session.add(batch)
    session.commit()
    session.refresh(batch)

    return {
        "batch": batch.model_dump(mode="json"),
        "imported_rows": snapshots + daily_metrics,
        "skipped_rows": max(0, total_rows - snapshots - daily_metrics),
        "warnings": warnings,
        "errors": errors,
        "created_posts": created_posts,
        "created_snapshots": snapshots,
        "created_daily_metrics": daily_metrics,
    }


def match_video_asset(session: Session, title: str | None, external_id: str | None = None) -> VideoAsset | None:
    if external_id:
        needle = external_id.lower()
        direct = session.exec(select(VideoAsset).where(VideoAsset.external_youtube_id == external_id)).first()
        if direct:
            return direct
        draft_direct = session.exec(select(PostDraft).where(PostDraft.external_post_id == external_id)).first()
        if draft_direct:
            video = session.get(VideoAsset, draft_direct.video_asset_id)
            if video:
                return video
        for draft in session.exec(select(PostDraft)).all():
            meta = draft.metadata_json or {}
            flattened = str(meta).lower()
            if needle and needle in flattened:
                video = session.get(VideoAsset, draft.video_asset_id)
                if video:
                    return video
    if not title:
        return None
    normalized = normalize_text(title)
    best = None
    best_score = 0.0
    for video in session.exec(select(VideoAsset)).all():
        score = jaccard(normalized, normalize_text(video.working_title or ""))
        if score > best_score:
            best = video
            best_score = score
    return best if best_score >= 0.55 else None


def normalize_text(value: str) -> set[str]:
    return {part for part in ''.join(ch.lower() if ch.isalnum() else ' ' for ch in value).split() if len(part) > 2}


def jaccard(a: set[str], b: set[str]) -> float:
    if not a or not b:
        return 0.0
    return len(a & b) / len(a | b)


def apply_deltas(snapshot: AnalyticsPostSnapshot, previous: AnalyticsPostSnapshot | None) -> None:
    if not previous:
        snapshot.trend = "new"
        return
    snapshot.views_delta = safe_delta(snapshot.views, previous.views)
    if snapshot.views_delta is not None and previous.views:
        snapshot.views_delta_pct = snapshot.views_delta / previous.views * 100
    snapshot.watch_time_delta = safe_delta(snapshot.watch_time_hours, previous.watch_time_hours)
    snapshot.impressions_delta = safe_delta(snapshot.impressions, previous.impressions)
    snapshot.subscriber_delta_change = safe_delta(snapshot.subscribers_delta, previous.subscribers_delta)
    snapshot.ctr_change = safe_delta(snapshot.impression_ctr_pct, previous.impression_ctr_pct)
    snapshot.retention_proxy_change = safe_delta(snapshot.retention_proxy_pct, previous.retention_proxy_pct)
    current_at = snapshot.snapshot_at
    previous_at = previous.snapshot_at
    if current_at.tzinfo is None:
        current_at = current_at.replace(tzinfo=timezone.utc)
    if previous_at.tzinfo is None:
        previous_at = previous_at.replace(tzinfo=timezone.utc)
    days = max(0.01, (current_at - previous_at).total_seconds() / 86400)
    snapshot.days_since_previous_snapshot = days
    if snapshot.views_delta is not None:
        snapshot.views_velocity_per_day = snapshot.views_delta / days
        snapshot.trend = "rising" if snapshot.views_delta > 10 else "slowing" if snapshot.views_delta < 1 else "stable"


def latest_snapshots(session: Session, include_demo: bool = False) -> list[AnalyticsPostSnapshot]:
    stmt = select(AnalyticsPostSnapshot).order_by(AnalyticsPostSnapshot.snapshot_at.desc())
    if not include_demo:
        stmt = stmt.where(AnalyticsPostSnapshot.is_demo == False)  # noqa: E712
    rows = session.exec(stmt).all()
    seen = set()
    latest = []
    for row in rows:
        key = (row.provider, row.account_id, row.external_post_id)
        if key in seen:
            continue
        seen.add(key)
        latest.append(row)
    return latest


def analytics_overview(session: Session) -> dict[str, Any]:
    latest = latest_snapshots(session)
    daily = session.exec(select(AnalyticsDailyAggregate).order_by(AnalyticsDailyAggregate.metric_date)).all()
    batches = session.exec(select(AnalyticsImportBatch).order_by(AnalyticsImportBatch.imported_at.desc())).all()
    total_views = sum(s.views or 0 for s in latest)
    total_watch = sum(s.watch_time_hours or 0 for s in latest)
    subscriber_delta = sum(s.subscribers_delta or 0 for s in latest)
    avg_retention = average([s.retention_proxy_pct for s in latest])
    best_by_views = max(latest, key=lambda s: s.views or -1, default=None)
    best_by_subs = max(latest, key=lambda s: s.subscribers_per_1000_views or -999999, default=None)
    daily_points = [{"date": d.metric_date.date().isoformat(), "views": d.views or 0} for d in daily]
    return {
        "kpis": {
            "total_views_latest": total_views,
            "views_last_7_days": sum(p["views"] for p in daily_points[-7:]),
            "views_last_28_days": sum(p["views"] for p in daily_points[-28:]),
            "subscriber_delta": subscriber_delta,
            "total_watch_time_hours": round(total_watch, 3),
            "average_retention_proxy": avg_retention,
            "best_video": best_by_views.title if best_by_views else None,
            "best_video_subscriber_conversion": best_by_subs.title if best_by_subs else None,
        },
        "daily_views": daily_points,
        "videos": [snapshot_to_dict(s) for s in latest],
        "imports": [b.model_dump(mode="json") for b in batches],
    }


def average(values):
    filtered = [v for v in values if v is not None]
    if not filtered:
        return None
    return sum(filtered) / len(filtered)


def snapshot_to_dict(s: AnalyticsPostSnapshot) -> dict[str, Any]:
    data = s.model_dump(mode="json")
    data["retention_label"] = "unknown" if s.retention_proxy_pct is None else "strong" if s.retention_proxy_pct >= 45 else "weak"
    return data


def analytics_insights(session: Session) -> list[dict[str, Any]]:
    latest = latest_snapshots(session)
    if not latest:
        return [{"type": "empty", "severity": "info", "title": "Import YouTube CSV to start performance learning", "message": "No analytics snapshots are available yet."}]
    views_values = [s.views or 0 for s in latest]
    retention_values = [s.retention_proxy_pct for s in latest if s.retention_proxy_pct is not None]
    impressions_values = [s.impressions or 0 for s in latest]
    ctr_values = [s.impression_ctr_pct for s in latest if s.impression_ctr_pct is not None]
    subs_values = [s.subscribers_per_1000_views for s in latest if s.subscribers_per_1000_views is not None]
    views_median = median(views_values) if views_values else 0
    retention_median = median(retention_values) if retention_values else 0
    impressions_median = median(impressions_values) if impressions_values else 0
    ctr_median = median(ctr_values) if ctr_values else 0
    subs_median = median(subs_values) if subs_values else 0
    insights: list[dict[str, Any]] = []
    for s in latest:
        title = s.title or s.external_post_id
        if (s.views or 0) > views_median and (s.retention_proxy_pct or 0) > retention_median:
            insights.append({"type": "winner", "severity": "success", "title": "Winner pattern", "message": f"{title}: keep this topic/hook style.", "external_post_id": s.external_post_id})
        if (s.retention_proxy_pct or 0) > retention_median and (s.views or 0) <= views_median:
            insights.append({"type": "packaging_problem", "severity": "warning", "title": "Packaging opportunity", "message": f"{title}: viewers who watch stay longer, but distribution/click packaging looks weak.", "external_post_id": s.external_post_id})
        if (s.impressions or 0) > impressions_median and (s.impression_ctr_pct or 0) < ctr_median:
            insights.append({"type": "hook_problem", "severity": "warning", "title": "Hook/CTR problem", "message": f"{title}: high exposure but weak click-through. Rework title/thumbnail/opening promise.", "external_post_id": s.external_post_id})
        if (s.views or 0) > views_median and (s.retention_proxy_pct or 0) < retention_median:
            insights.append({"type": "retention_problem", "severity": "warning", "title": "Retention problem", "message": f"{title}: topic attracts views, but pacing or first seconds may lose viewers.", "external_post_id": s.external_post_id})
        if (s.subscribers_per_1000_views or -999999) > subs_median:
            insights.append({"type": "subscriber_converter", "severity": "success", "title": "Subscriber converter", "message": f"{title}: converts viewers into subscribers. Reuse a similar promise/CTA.", "external_post_id": s.external_post_id})
        if s.trend == "slowing":
            insights.append({"type": "negative_trend", "severity": "info", "title": "Distribution slowing", "message": f"{title}: distribution flattened since previous import. Consider follow-up or repackaging.", "external_post_id": s.external_post_id})
    insights.extend(topic_opportunities(session, latest))
    return insights[:12]


def topic_opportunities(session: Session, latest: list[AnalyticsPostSnapshot]) -> list[dict[str, Any]]:
    groups: dict[str, list[AnalyticsPostSnapshot]] = defaultdict(list)
    posts = {p.external_post_id: p for p in session.exec(select(ExternalPost)).all()}
    for s in latest:
        post = posts.get(s.external_post_id)
        if post and post.topic_id:
            groups[post.topic_id].append(s)
    all_avg = average([s.views for s in latest]) or 0
    out = []
    for topic_id, rows in groups.items():
        topic_avg = average([s.views for s in rows]) or 0
        if len(rows) < 3 and topic_avg > all_avg:
            topic = session.get(Theme, topic_id)
            out.append({"type": "topic_opportunity", "severity": "success", "title": "Topic opportunity", "message": f"{topic.title if topic else topic_id}: few videos but above-average views. Make 3 more variants before switching topic.", "topic_id": topic_id})
    return out
