from datetime import datetime
from typing import List, Optional

from pymongo import ASCENDING, UpdateOne

from app.core.database import get_db
from app.models.embedding_job import EmbeddingJob


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


async def upsert_jobs(jobs: List[EmbeddingJob]) -> dict:
    if not jobs:
        return {"requested": 0, "upserted": 0}

    await _ensure_indexes()
    db = get_db()

    operations = []
    for job in jobs:
        payload = job.model_dump(exclude={"id"})
        now = datetime.utcnow()
        payload["created_at"] = now
        payload.pop("updated_at", None)
        operations.append(
            UpdateOne(
                {"dedupe_key": job.dedupe_key},
                {
                    "$setOnInsert": payload,
                    "$set": {"updated_at": now},
                },
                upsert=True,
            )
        )

    result = await db.embedding_jobs.bulk_write(operations, ordered=False)
    return {
        "requested": len(jobs),
        "upserted": result.upserted_count,
    }


async def get_jobs_by_status(
    statuses: List[str],
    limit: int = 100,
    ready_only: bool = False,
) -> List[EmbeddingJob]:
    await _ensure_indexes()
    db = get_db()
    query: dict = {"status": {"$in": statuses}}
    if ready_only:
        query["$or"] = [
            {"next_retry_at": {"$exists": False}},
            {"next_retry_at": None},
            {"next_retry_at": {"$lte": datetime.utcnow()}},
        ]

    cursor = db.embedding_jobs.find(query).sort("created_at", ASCENDING).limit(limit)
    return [EmbeddingJob(**_serialize(doc)) async for doc in cursor]


async def get_job(job_id: str) -> Optional[EmbeddingJob]:
    await _ensure_indexes()
    db = get_db()
    doc = await db.embedding_jobs.find_one({"job_id": job_id})
    return EmbeddingJob(**_serialize(doc)) if doc else None


async def update_job_state(
    job_id: str,
    status: str,
    *,
    last_error: Optional[str] = None,
    processing_started: bool = False,
    processing_completed: bool = False,
    published: bool = False,
    increment_retry: bool = False,
    retry_event: Optional[dict] = None,
    next_retry_at: Optional[datetime] = None,
    clear_next_retry: bool = False,
) -> None:
    await _ensure_indexes()
    db = get_db()

    now = datetime.utcnow()
    set_values = {
        "status": status,
        "updated_at": now,
    }

    if last_error is not None:
        set_values["last_error"] = last_error

    if processing_started:
        set_values["processing_started_at"] = now

    if processing_completed:
        set_values["processing_completed_at"] = now

    if published:
        set_values["published_at"] = now

    if next_retry_at is not None:
        set_values["next_retry_at"] = next_retry_at

    if clear_next_retry:
        set_values["next_retry_at"] = None

    update_doc = {"$set": set_values}

    if increment_retry:
        update_doc.setdefault("$inc", {})["retry_count"] = 1

    if retry_event is not None:
        event = dict(retry_event)
        event.setdefault("at", now.isoformat())
        update_doc.setdefault("$push", {})["retry_history"] = event

    await db.embedding_jobs.update_one({"job_id": job_id}, update_doc)
