from __future__ import annotations

from datetime import datetime
from typing import List, Optional

from pymongo import ASCENDING, ReturnDocument, UpdateOne

from app.core.database import get_db
from app.models.embedding_run import EmbeddingRun, build_run_dedupe_key, build_run_id
from app.models.enums import EmbedStatus


_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()
    collection = db.embedding_runs
    await collection.create_index([("run_id", ASCENDING)], unique=True)
    await collection.create_index([("dedupe_key", ASCENDING)], unique=True)
    await collection.create_index([("repo_id", ASCENDING), ("snapshot_id", ASCENDING)])
    await collection.create_index([("status", ASCENDING), ("updated_at", ASCENDING)])
    _INDEXES_READY = True


async def upsert_run(
    *,
    repo_id: str,
    snapshot_id: str,
    event_id: str,
    run_type: str = "snapshot_batch",
    status: EmbedStatus = EmbedStatus.PENDING,
    celery_task_id: Optional[str] = None,
) -> EmbeddingRun:
    await _ensure_indexes()
    db = get_db()

    dedupe_key = build_run_dedupe_key(repo_id, snapshot_id, event_id)
    run_id = build_run_id(dedupe_key)
    now = datetime.utcnow()

    payload = {
        "run_id": run_id,
        "dedupe_key": dedupe_key,
        "repo_id": repo_id,
        "snapshot_id": snapshot_id,
        "event_id": event_id,
        "run_type": run_type,
        "status": status.value,
        "updated_at": now,
    }
    if celery_task_id is not None:
        payload["celery_task_id"] = celery_task_id

    doc = await db.embedding_runs.find_one_and_update(
        {"dedupe_key": dedupe_key},
        {
            "$set": payload,
            "$setOnInsert": {
                "created_at": now,
                "total_records": 0,
                "embedded_records": 0,
                "failed_records": 0,
                "retry_count": 0,
            },
        },
        upsert=True,
        return_document=ReturnDocument.AFTER,
    )
    return EmbeddingRun.model_validate(_serialize(doc))


async def get_run(run_id: str) -> Optional[EmbeddingRun]:
    await _ensure_indexes()
    db = get_db()
    doc = await db.embedding_runs.find_one({"run_id": run_id})
    if not doc:
        return None
    return EmbeddingRun.model_validate(_serialize(doc))


async def get_run_by_dedupe_key(dedupe_key: str) -> Optional[EmbeddingRun]:
    await _ensure_indexes()
    db = get_db()
    doc = await db.embedding_runs.find_one({"dedupe_key": dedupe_key})
    if not doc:
        return None
    return EmbeddingRun.model_validate(_serialize(doc))


async def update_run_state(
    run_id: str,
    status: EmbedStatus,
    *,
    total_records: Optional[int] = None,
    embedded_records: Optional[int] = None,
    failed_records: Optional[int] = None,
    last_error: Optional[str] = None,
    increment_retry: bool = False,
    mark_started: bool = False,
    mark_completed: bool = False,
    celery_task_id: Optional[str] = None,
) -> bool:
    await _ensure_indexes()
    db = get_db()
    now = datetime.utcnow()

    set_values: dict = {
        "status": status.value,
        "updated_at": now,
    }
    if total_records is not None:
        set_values["total_records"] = total_records
    if embedded_records is not None:
        set_values["embedded_records"] = embedded_records
    if failed_records is not None:
        set_values["failed_records"] = failed_records
    if last_error is not None:
        set_values["last_error"] = last_error
    if celery_task_id is not None:
        set_values["celery_task_id"] = celery_task_id
    if mark_started:
        set_values["started_at"] = now
    if mark_completed:
        set_values["completed_at"] = now

    update_doc: dict = {"$set": set_values}
    if increment_retry:
        update_doc["$inc"] = {"retry_count": 1}

    result = await db.embedding_runs.update_one({"run_id": run_id}, update_doc)
    return result.matched_count > 0


async def list_runs_by_snapshot(
    repo_id: str,
    snapshot_id: str,
    *,
    statuses: Optional[List[EmbedStatus]] = None,
    limit: int = 50,
) -> List[EmbeddingRun]:
    await _ensure_indexes()
    db = get_db()
    query: dict = {"repo_id": repo_id, "snapshot_id": snapshot_id}
    if statuses:
        query["status"] = {"$in": [status.value for status in statuses]}

    cursor = (
        db.embedding_runs.find(query).sort("created_at", ASCENDING).limit(limit)
    )
    return [EmbeddingRun.model_validate(_serialize(doc)) async for doc in cursor]
