from datetime import datetime
from typing import Optional

from pymongo import ASCENDING, ReturnDocument

from app.core.database import get_db
from app.models.parse_run import DocParseRun, ParseRunStatus


_INDEXES_READY = False


def _serialize(doc: dict) -> dict:
    if doc and "_id" in doc:
        doc["id"] = str(doc.pop("_id"))
    return doc


async def _ensure_indexes() -> None:
    global _INDEXES_READY
    if _INDEXES_READY:
        return

    db = get_db()
    await db.doc_parse_runs.create_index([("run_id", ASCENDING)], unique=True)
    await db.doc_parse_runs.create_index(
        [("dedupe_key", ASCENDING)],
        unique=True,
        sparse=True,
    )
    await db.doc_parse_runs.create_index(
        [("repo_id", ASCENDING), ("snapshot_id", ASCENDING), ("doc_path", ASCENDING)]
    )
    _INDEXES_READY = True


async def create_run(run: DocParseRun) -> str:
    await _ensure_indexes()
    db = get_db()
    result = await db.doc_parse_runs.insert_one(run.model_dump(exclude={"id"}))
    return str(result.inserted_id)


async def upsert_run_for_event(run: DocParseRun) -> DocParseRun:
    await _ensure_indexes()
    db = get_db()
    payload = run.model_dump(exclude={"id"})
    payload.pop("created_at", None)
    result = await db.doc_parse_runs.find_one_and_update(
        {"dedupe_key": run.dedupe_key},
        {"$set": payload, "$setOnInsert": {"created_at": datetime.utcnow()}},
        upsert=True,
        return_document=ReturnDocument.AFTER,
    )
    return DocParseRun(**_serialize(result))


async def get_completed_run(dedupe_key: str) -> Optional[DocParseRun]:
    await _ensure_indexes()
    db = get_db()
    doc = await db.doc_parse_runs.find_one(
        {"dedupe_key": dedupe_key, "status": ParseRunStatus.COMPLETED}
    )
    return DocParseRun(**_serialize(doc)) if doc else None


async def get_run_by_id(run_id: str) -> Optional[DocParseRun]:
    await _ensure_indexes()
    db = get_db()
    doc = await db.doc_parse_runs.find_one({"run_id": run_id})
    return DocParseRun(**_serialize(doc)) if doc else None


async def update_run_status(
    run_id: str,
    status: str,
    chunks_created: int = 0,
    error: Optional[str] = None,
) -> None:
    await _ensure_indexes()
    db = get_db()
    update: dict = {"status": status, "chunks_created": chunks_created}
    if status == ParseRunStatus.PROCESSING:
        update["started_at"] = datetime.utcnow()
    if status in (ParseRunStatus.COMPLETED, ParseRunStatus.FAILED):
        update["completed_at"] = datetime.utcnow()
    if error:
        update["error"] = error
        update["last_error"] = error
    if status in (ParseRunStatus.COMPLETED, ParseRunStatus.FAILED):
        update["next_retry_at"] = None
    await db.doc_parse_runs.update_one({"run_id": run_id}, {"$set": update})


async def mark_run_queued(run_id: str) -> None:
    db = get_db()
    await db.doc_parse_runs.update_one(
        {"run_id": run_id},
        {
            "$set": {
                "status": ParseRunStatus.QUEUED,
                "error": None,
            }
        },
    )


async def mark_run_retrying(
    run_id: str,
    retry_count: int,
    error: str,
    next_retry_at: datetime,
) -> None:
    db = get_db()
    await db.doc_parse_runs.update_one(
        {"run_id": run_id},
        {
            "$set": {
                "status": ParseRunStatus.RETRYING,
                "retry_count": retry_count,
                "last_error": error,
                "error": error,
                "next_retry_at": next_retry_at,
            }
        },
    )


async def claim_run_for_processing(
    run_id: str,
    stale_after_seconds: int,
) -> bool:
    db = get_db()
    now = datetime.utcnow()
    stale_cutoff = datetime.utcfromtimestamp(now.timestamp() - max(1, stale_after_seconds))

    result = await db.doc_parse_runs.update_one(
        {
            "run_id": run_id,
            "$or": [
                {"status": {"$in": [ParseRunStatus.PENDING, ParseRunStatus.QUEUED, ParseRunStatus.RETRYING]}},
                {
                    "status": ParseRunStatus.PROCESSING,
                    "started_at": {"$lte": stale_cutoff},
                },
            ],
        },
        {
            "$set": {
                "status": ParseRunStatus.PROCESSING,
                "started_at": now,
                "error": None,
                "next_retry_at": None,
            }
        },
    )
    return result.modified_count == 1
