from __future__ import annotations

from typing import Any, Optional

from app.core.database import get_code_chunks_db
from app.core.logger import logger
from app.models.upstream.snapshot_graph import SnapshotGraphDocument

SNAPSHOT_GRAPH_PARTS_COLLECTION = "snapshot_graph_parts"


def merged_artifact_id(snapshot_id: str) -> str:
    return f"art_{snapshot_id}_merged"


def is_chunked_manifest(doc: dict[str, Any]) -> bool:
    if doc.get("chunked") is True:
        return True
    part_count = int(doc.get("part_count") or 0)
    has_inline_payload = bool(doc.get("nodes")) or bool(doc.get("edges"))
    return part_count > 0 and not has_inline_payload


def reassemble_graph_payload(parts: list[dict[str, Any]]) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]:
    nodes: list[dict[str, Any]] = []
    edges: list[dict[str, Any]] = []
    for part in sorted(parts, key=lambda item: int(item.get("part_index", 0))):
        nodes.extend(part.get("nodes") or [])
        edges.extend(part.get("edges") or [])
    return nodes, edges


async def _load_snapshot_manifest(
    db: Any,
    *,
    repo_id: str,
    snapshot_id: str,
    artifact_id: str,
) -> Optional[dict[str, Any]]:
    doc = await db.snapshot_graphs.find_one({"_id": artifact_id})
    if doc is None:
        doc = await db.snapshot_graphs.find_one(
            {"artifact_id": artifact_id, "repo_id": repo_id, "snapshot_id": snapshot_id}
        )
    if doc is None:
        doc = await db.snapshot_graphs.find_one({"repo_id": repo_id, "snapshot_id": snapshot_id})
    return doc


async def _load_graph_parts(db: Any, artifact_id: str, part_count: int) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]:
    parts_collection = getattr(db, SNAPSHOT_GRAPH_PARTS_COLLECTION, None)
    if parts_collection is None:
        logger.warning(
            "snapshot_graph_parts collection unavailable for artifact_id=%s",
            artifact_id,
        )
        return [], []

    cursor = parts_collection.find({"artifact_id": artifact_id})
    parts: list[dict[str, Any]] = []
    async for part_doc in cursor:
        parts.append(part_doc)

    if part_count > 0 and len(parts) != part_count:
        logger.warning(
            "snapshot graph part count mismatch for artifact_id=%s expected=%d found=%d",
            artifact_id,
            part_count,
            len(parts),
        )

    return reassemble_graph_payload(parts)


async def get_merged(repo_id: str, snapshot_id: str) -> Optional[SnapshotGraphDocument]:
    """Load the merged snapshot graph for a repo snapshot."""
    db = get_code_chunks_db()
    artifact_id = merged_artifact_id(snapshot_id)

    doc = await _load_snapshot_manifest(
        db,
        repo_id=repo_id,
        snapshot_id=snapshot_id,
        artifact_id=artifact_id,
    )
    if doc is None:
        return None

    payload = dict(doc)
    if is_chunked_manifest(payload):
        resolved_artifact_id = str(payload.get("artifact_id") or payload.get("_id") or artifact_id)
        part_count = int(payload.get("part_count") or 0)
        nodes, edges = await _load_graph_parts(db, resolved_artifact_id, part_count)
        payload["nodes"] = nodes
        payload["edges"] = edges

    graph = SnapshotGraphDocument.from_mongo(payload)
    if graph.repo_id != repo_id or graph.snapshot_id != snapshot_id:
        return None
    return graph
