from __future__ import annotations

from datetime import datetime, timezone
from hashlib import sha256

from app.core.config import settings
from app.core.database import get_indexing_runs_db
from app.core.logger import logger
from app.models.enums import EmbedStatus, SourceType
from app.repositories import embedding_record_repo
from app.repositories.upstream import commit_fanin_repo
from app.services.commit_fanin import resolve_expected_commit_deltas

STAGE_RUNNING = "running"
STAGE_COMPLETED = "completed"
STAGE_FAILED = "failed"
EMBEDDING_FAILURE_THRESHOLD = 0.05


def _enabled() -> bool:
    return settings.INDEXING_RUNS_ENABLED


async def mark_embedding_running(repo_id: str, snapshot_id: str) -> None:
    """Mark the embedding stage as running for a snapshot-scoped indexing run."""
    if not _enabled():
        return

    try:
        db = get_indexing_runs_db()
        existing = await db.indexing_runs.find_one(
            {"repo_id": repo_id, "snapshot_id": snapshot_id},
            projection={"stages.embedding": 1},
        )
        if not existing:
            return

        current = (existing.get("stages") or {}).get("embedding")
        if current == STAGE_COMPLETED:
            return

        await db.indexing_runs.update_one(
            {"repo_id": repo_id, "snapshot_id": snapshot_id},
            {"$set": {"stages.embedding": STAGE_RUNNING}},
        )
    except Exception as exc:
        logger.warning(
            "Failed to mark embedding stage running repo_id={repo_id} snapshot_id={snapshot_id} error={error}",
            repo_id=repo_id,
            snapshot_id=snapshot_id,
            error=str(exc),
        )


async def mark_embedding_failed(repo_id: str, snapshot_id: str) -> None:
    """Mark the embedding stage as failed (snapshot batch failure)."""
    if not _enabled():
        return

    try:
        db = get_indexing_runs_db()
        result = await db.indexing_runs.update_one(
            {"repo_id": repo_id, "snapshot_id": snapshot_id},
            {"$set": {"stages.embedding": STAGE_FAILED}},
        )
        if result.matched_count == 0:
            logger.debug(
                "No indexing run to mark embedding failed repo_id={repo_id} snapshot_id={snapshot_id}",
                repo_id=repo_id,
                snapshot_id=snapshot_id,
            )
    except Exception as exc:
        logger.warning(
            "Failed to mark embedding stage failed repo_id={repo_id} snapshot_id={snapshot_id} error={error}",
            repo_id=repo_id,
            snapshot_id=snapshot_id,
            error=str(exc),
        )


async def increment_failed_embeddings(
    repo_id: str,
    snapshot_id: str,
    *,
    source_type: str,
    source_id: str,
    event_id: str,
    message: str,
    count: int = 1,
) -> None:
    """Record permanent embedding failures and persist a diagnostic."""
    if not _enabled():
        return

    try:
        db = get_indexing_runs_db()
        dedupe_key = f"{source_type}:{source_id}"
        result = await db.indexing_runs.update_one(
            {
                "repo_id": repo_id,
                "snapshot_id": snapshot_id,
                "failed_embedding_sources": {"$ne": dedupe_key},
            },
            {
                "$addToSet": {"failed_embedding_sources": dedupe_key},
                "$inc": {"failed_embeddings": max(1, count)},
            },
        )
        await _sync_embedding_failure_rate(repo_id, snapshot_id)
        if result.matched_count:
            await save_indexing_diagnostic(
                repo_id,
                snapshot_id,
                stage="embedding",
                source_type=source_type,
                source_id=source_id,
                event_id=event_id,
                code="embedding_failed",
                message=message,
            )
    except Exception as exc:
        logger.warning(
            "Failed to record embedding failure repo_id={repo_id} snapshot_id={snapshot_id} source_id={source_id} error={error}",
            repo_id=repo_id,
            snapshot_id=snapshot_id,
            source_id=source_id,
            error=str(exc),
        )


async def save_indexing_diagnostic(
    repo_id: str,
    snapshot_id: str,
    *,
    stage: str,
    source_type: str,
    source_id: str,
    event_id: str,
    code: str,
    message: str,
    severity: str = "error",
) -> None:
    if not _enabled():
        return

    db = get_indexing_runs_db()
    diagnostic_id = _diagnostic_id(repo_id, snapshot_id, stage, source_type, source_id, event_id, code)
    await db.indexing_diagnostics.update_one(
        {"diagnostic_id": diagnostic_id},
        {
            "$set": {
                "diagnostic_id": diagnostic_id,
                "repo_id": repo_id,
                "snapshot_id": snapshot_id,
                "stage": stage,
                "source_type": source_type,
                "source_id": source_id,
                "event_id": event_id,
                "code": code,
                "severity": severity,
                "message": message,
                "created_at": datetime.now(timezone.utc),
            }
        },
        upsert=True,
    )


async def _sync_embedding_failure_rate(repo_id: str, snapshot_id: str) -> None:
    db = get_indexing_runs_db()
    run = await db.indexing_runs.find_one(
        {"repo_id": repo_id, "snapshot_id": snapshot_id},
        projection={"failed_embeddings": 1},
    )
    if not run:
        return

    failed = int(run.get("failed_embeddings") or 0)
    embedded = await embedding_record_repo.count_records_by_snapshot(
        repo_id,
        snapshot_id,
        embed_status=EmbedStatus.EMBEDDED,
    )
    denominator = embedded + failed
    failure_rate = (failed / denominator) if denominator > 0 else 0
    set_values = {"embedding_failure_rate": failure_rate}
    if failure_rate >= EMBEDDING_FAILURE_THRESHOLD:
        set_values["stages.embedding"] = STAGE_FAILED
    await db.indexing_runs.update_one(
        {"repo_id": repo_id, "snapshot_id": snapshot_id},
        {"$set": set_values},
    )


def _diagnostic_id(*parts: str) -> str:
    joined = "|".join(parts)
    return "diag_" + sha256(joined.encode("utf-8")).hexdigest()[:16]


async def maybe_mark_embedding_completed(repo_id: str, snapshot_id: str) -> None:
    """
    Mark embedding completed when snapshot chunk embedding is done and all delta-backed
    commit summaries are embedded.
    """
    if not _enabled():
        return

    try:
        db = get_indexing_runs_db()
        run = await db.indexing_runs.find_one(
            {"repo_id": repo_id, "snapshot_id": snapshot_id},
            projection={
                "expected_commits": 1,
                "expected_commit_deltas": 1,
                "stages.embedding": 1,
            },
        )
        if not run:
            return

        expected_commits = int(run.get("expected_commits") or 0)
        pinned_deltas = int(run.get("expected_commit_deltas") or 0)
        if expected_commits <= 0 and pinned_deltas <= 0:
            await db.indexing_runs.update_one(
                {"repo_id": repo_id, "snapshot_id": snapshot_id},
                {"$set": {"stages.embedding": STAGE_COMPLETED}},
            )
            return

        completed_commits = await commit_fanin_repo.count_completed_commits(repo_id, snapshot_id)
        delta_count = await commit_fanin_repo.count_graph_deltas(repo_id, snapshot_id)
        target_commits = resolve_expected_commit_deltas(
            expected_commits,
            pinned_deltas,
            completed_commits,
            delta_count,
        )

        embedded_commits = await embedding_record_repo.count_records_by_snapshot(
            repo_id,
            snapshot_id,
            source_type=SourceType.COMMIT,
            embed_status=EmbedStatus.EMBEDDED,
        )
        should_complete = False
        if target_commits > 0:
            should_complete = embedded_commits >= target_commits
        elif completed_commits >= expected_commits and delta_count == 0:
            # All commits processed with no delta-backed summaries to embed.
            should_complete = True

        if should_complete:
            await db.indexing_runs.update_one(
                {"repo_id": repo_id, "snapshot_id": snapshot_id},
                {"$set": {"stages.embedding": STAGE_COMPLETED}},
            )
    except Exception as exc:
        logger.warning(
            "Failed to update embedding stage completion repo_id={repo_id} snapshot_id={snapshot_id} error={error}",
            repo_id=repo_id,
            snapshot_id=snapshot_id,
            error=str(exc),
        )
