from __future__ import annotations

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_record import EmbeddingRecord
from app.models.embedding_record_upsert import EmbeddingRecordUpsert
from app.models.enums import EmbedStatus, SourceType
from app.models.stored_embedding_record import StoredEmbeddingRecord, build_record_dedupe_key


_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_records
    await collection.create_index([("record_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(
        [("repo_id", ASCENDING), ("snapshot_id", ASCENDING), ("source_type", ASCENDING)]
    )
    await collection.create_index([("embed_status", ASCENDING), ("updated_at", ASCENDING)])
    _INDEXES_READY = True


def _record_to_document(
    record: EmbeddingRecord,
    *,
    embed_status: EmbedStatus = EmbedStatus.PENDING,
    upstream_chunk_id: Optional[str] = None,
    upstream_analysis_id: Optional[str] = None,
) -> dict:
    stored = StoredEmbeddingRecord(
        **record.model_dump(),
        dedupe_key=build_record_dedupe_key(
            record.repo_id, record.snapshot_id, record.record_id
        ),
        embed_status=embed_status,
        upstream_chunk_id=upstream_chunk_id,
        upstream_analysis_id=upstream_analysis_id,
    )
    return stored.model_dump(mode="json", exclude={"id"})


def _upsert_item_to_operation(item: EmbeddingRecordUpsert, now: datetime) -> UpdateOne:
    document = _record_to_document(
        item.record,
        embed_status=item.embed_status,
        upstream_chunk_id=item.upstream_chunk_id,
        upstream_analysis_id=item.upstream_analysis_id,
    )
    document["updated_at"] = now
    # created_at must only appear in $setOnInsert — StoredEmbeddingRecord defaults
    # would otherwise land in $set and conflict with MongoDB error code 40.
    document.pop("created_at", None)
    return UpdateOne(
        {"record_id": item.record.record_id},
        {
            "$set": document,
            "$setOnInsert": {"created_at": now},
        },
        upsert=True,
    )


async def upsert_record_items(items: List[EmbeddingRecordUpsert]) -> dict:
    if not items:
        return {"requested": 0, "upserted": 0, "modified": 0}

    await _ensure_indexes()
    db = get_db()
    now = datetime.utcnow()
    operations = [_upsert_item_to_operation(item, now) for item in items]
    result = await db.embedding_records.bulk_write(operations, ordered=False)
    return {
        "requested": len(items),
        "upserted": result.upserted_count,
        "modified": result.modified_count,
    }


async def upsert_records(
    records: List[EmbeddingRecord],
    *,
    embed_status: EmbedStatus = EmbedStatus.PENDING,
    upstream_chunk_id: Optional[str] = None,
    upstream_analysis_id: Optional[str] = None,
) -> dict:
    if not records:
        return {"requested": 0, "upserted": 0, "modified": 0}

    items = [
        EmbeddingRecordUpsert(
            record=record,
            embed_status=embed_status,
            upstream_chunk_id=upstream_chunk_id,
            upstream_analysis_id=upstream_analysis_id,
        )
        for record in records
    ]
    return await upsert_record_items(items)


async def get_record(record_id: str) -> Optional[StoredEmbeddingRecord]:
    await _ensure_indexes()
    db = get_db()
    doc = await db.embedding_records.find_one({"record_id": record_id})
    if not doc:
        return None
    return StoredEmbeddingRecord.model_validate(_serialize(doc))


async def get_records_by_ids(record_ids: List[str]) -> dict[str, StoredEmbeddingRecord]:
    """Batch-load embedding records keyed by record_id."""
    if not record_ids:
        return {}

    await _ensure_indexes()
    db = get_db()
    unique_ids = list(dict.fromkeys(record_ids))
    cursor = db.embedding_records.find({"record_id": {"$in": unique_ids}})
    records: dict[str, StoredEmbeddingRecord] = {}
    async for doc in cursor:
        record = StoredEmbeddingRecord.model_validate(_serialize(doc))
        records[record.record_id] = record
    return records


async def list_records_by_snapshot(
    repo_id: str,
    snapshot_id: str,
    *,
    source_type: Optional[SourceType] = None,
    embed_status: Optional[EmbedStatus] = None,
    skip: int = 0,
    limit: int = 500,
) -> List[StoredEmbeddingRecord]:
    await _ensure_indexes()
    db = get_db()
    query: dict = {"repo_id": repo_id, "snapshot_id": snapshot_id}
    if source_type is not None:
        query["source_type"] = source_type.value
    if embed_status is not None:
        query["embed_status"] = embed_status.value

    cursor = (
        db.embedding_records.find(query).skip(skip).limit(limit).sort("record_id", ASCENDING)
    )
    return [StoredEmbeddingRecord.model_validate(_serialize(doc)) async for doc in cursor]


async def count_records_by_snapshot(
    repo_id: str,
    snapshot_id: str,
    *,
    source_type: Optional[SourceType] = None,
    embed_status: Optional[EmbedStatus] = None,
) -> int:
    await _ensure_indexes()
    db = get_db()
    query: dict = {"repo_id": repo_id, "snapshot_id": snapshot_id}
    if source_type is not None:
        query["source_type"] = source_type.value
    if embed_status is not None:
        query["embed_status"] = embed_status.value
    return await db.embedding_records.count_documents(query)


async def update_embed_status(
    record_id: str,
    status: EmbedStatus,
    *,
    qdrant_point_id: Optional[str] = None,
    last_error: Optional[str] = None,
    clear_error: bool = False,
    clear_qdrant_point_id: bool = False,
) -> bool:
    await _ensure_indexes()
    db = get_db()
    now = datetime.utcnow()

    set_values: dict = {
        "embed_status": status.value,
        "updated_at": now,
    }
    if qdrant_point_id is not None:
        set_values["qdrant_point_id"] = qdrant_point_id
    if clear_qdrant_point_id:
        set_values["qdrant_point_id"] = None
    if last_error is not None:
        set_values["last_error"] = last_error
    if clear_error:
        set_values["last_error"] = None
    if status == EmbedStatus.EMBEDDED:
        set_values["embedded_at"] = now

    result = await db.embedding_records.update_one(
        {"record_id": record_id},
        {"$set": set_values},
    )
    return result.matched_count > 0


async def delete_records_by_ids(record_ids: List[str]) -> int:
    if not record_ids:
        return 0
    await _ensure_indexes()
    db = get_db()
    result = await db.embedding_records.delete_many({"record_id": {"$in": record_ids}})
    return result.deleted_count
